import argparse import json import os import time from pathlib import Path import torch from peft import PeftModel from transformers import AutoModelForCausalLM, AutoTokenizer from helpfulness_scheme_interface import ( _compute_allowed_ids_scheme_a, _compute_allowed_ids_scheme_b_temp, ) PROMPT_BEGIN = "BEGINNING OF CONVERSATION: " PROMPT_USER = "USER: {input} " PROMPT_ASSISTANT = "ASSISTANT:" PROMPT_INPUT_ALPACA = PROMPT_BEGIN + PROMPT_USER + PROMPT_ASSISTANT def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--prompt_file", type=str, required=True) parser.add_argument("--output_dir", type=str, required=True) parser.add_argument("--model_base_name_or_path", type=str, required=True) parser.add_argument("--model_arm_helpfulness_name_or_path", type=str, required=True) parser.add_argument("--model_arm_harmlessness_name_or_path", type=str, required=True) parser.add_argument("--alpha_helpfulness", type=float, required=True) parser.add_argument("--alpha_harmlessness", type=float, required=True) parser.add_argument("--model_name_for_logging", type=str, required=True) parser.add_argument("--scheme", choices=["a", "b_temp_support"], required=True) parser.add_argument("--base_prob_threshold", type=float, default=0.0008) parser.add_argument("--support_temperature", type=float, default=10.0) parser.add_argument("--max_new_tokens", type=int, default=512) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--resume", action="store_true") return parser.parse_args() def load_models(args: argparse.Namespace): tokenizer = AutoTokenizer.from_pretrained(args.model_base_name_or_path) if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token base_model = AutoModelForCausalLM.from_pretrained( args.model_base_name_or_path, torch_dtype=torch.bfloat16, device_map="auto", trust_remote_code=True, ) model = PeftModel.from_pretrained( base_model, args.model_arm_helpfulness_name_or_path, adapter_name="helpfulness", ) model.load_adapter(args.model_arm_harmlessness_name_or_path, adapter_name="harmlessness") model.eval() return tokenizer, model def forward_last_logprobs(model, input_ids: list[int], mode: str) -> torch.Tensor: device = next(model.parameters()).device tensor = torch.tensor([input_ids], device=device) with torch.no_grad(): if mode == "base": with model.disable_adapter(): logits = model(input_ids=tensor, use_cache=False).logits[0, -1].float().cpu() else: model.set_adapter(mode) logits = model(input_ids=tensor, use_cache=False).logits[0, -1].float().cpu() return torch.log_softmax(logits, dim=-1) def sample_response( prompt: str, tokenizer, model, args: argparse.Namespace, ) -> tuple[str, list[dict]]: prompt_ids = tokenizer(prompt, add_special_tokens=False)["input_ids"] response_ids: list[int] = [] trace: list[dict] = [] for step_idx in range(args.max_new_tokens): full_prefix = prompt_ids + response_ids logp_base = forward_last_logprobs(model, full_prefix, "base") logp_help = forward_last_logprobs(model, full_prefix, "helpfulness") logp_harm = forward_last_logprobs(model, full_prefix, "harmlessness") prob_base = torch.exp(logp_base) if args.scheme == "a": allowed_ids = _compute_allowed_ids_scheme_a(prob_base, args.base_prob_threshold) support_temperature = None else: allowed_ids, _ = _compute_allowed_ids_scheme_b_temp(logp_base, args.support_temperature) support_temperature = args.support_temperature score_s = ( logp_base + args.alpha_helpfulness * logp_help + args.alpha_harmlessness * logp_harm ) filtered_score = score_s[allowed_ids] filtered_probs = torch.softmax(filtered_score, dim=-1) sampled_index = int(torch.multinomial(filtered_probs, 1).item()) token_id = int(allowed_ids[sampled_index].item()) trace.append( { "step_index": step_idx + 1, "token_id": token_id, "token_text": tokenizer.decode([token_id]), "allowed_token_count": int(allowed_ids.numel()), "scheme": args.scheme, "base_prob_threshold": args.base_prob_threshold if args.scheme == "a" else None, "support_temperature": support_temperature, } ) if token_id == tokenizer.eos_token_id: break response_ids.append(token_id) return tokenizer.decode(response_ids, skip_special_tokens=True), trace def main() -> None: args = parse_args() torch.manual_seed(args.seed) output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) out_path = output_dir / "generation.json" if out_path.exists() and not args.resume: raise SystemExit(f"{out_path} exists; pass --resume to continue.") with open(args.prompt_file, "r", encoding="utf-8") as handle: prompts = json.load(handle) tokenizer, model = load_models(args) if args.resume and out_path.exists(): with open(out_path, "r", encoding="utf-8") as handle: outputs = json.load(handle) start = len(outputs) else: outputs = [] start = 0 for idx in range(start, len(prompts)): row = prompts[idx] formatted_prompt = PROMPT_INPUT_ALPACA.format(input=row["prompt"]) tic = time.time() response, trace = sample_response(formatted_prompt, tokenizer, model, args) elapsed = time.time() - tic outputs.append( { "uid": row["uid"], "prompt": row["prompt"], "response": response, "model": args.model_name_for_logging, "elapsed": elapsed, "scheme": args.scheme, "alpha_helpfulness": args.alpha_helpfulness, "alpha_harmlessness": args.alpha_harmlessness, "trace_first_10_steps": trace[:10], } ) if idx % 3 == 1: with open(out_path, "w", encoding="utf-8") as handle: json.dump(outputs, handle, ensure_ascii=False, indent=2) with open(out_path, "w", encoding="utf-8") as handle: json.dump(outputs, handle, ensure_ascii=False, indent=2) if __name__ == "__main__": main()