| 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() |
|
|