genarm-h0p01-c0175-eval-code / scripts /generate_old_genarm_custom_scheme.py
sheng22213's picture
Add GenARM h0p01 c=0.175 eval code
33b38c9 verified
Raw
History Blame Contribute Delete
6.67 kB
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()