mypo-training / examples /reproduce_v3.py
joshuasundance's picture
Clarify reproduce_v3 as smoke test, not exact eval-row replay
9b0cc8f verified
Raw
History Blame Contribute Delete
4.63 kB
#!/usr/bin/env python3
"""Minimal single-prompt reproducibility demo for MyPO v3.
Loads the four relevant models and prints side-by-side generations for a prompt:
- base: Qwen/Qwen2.5-Coder-1.5B-Instruct
- sft: joshuasundance/mypo-qwen2.5-coder-1.5b-sft
- dpo-v2: joshuasundance/mypo-qwen2.5-coder-1.5b-dpo-v2
- dpo-v3: joshuasundance/mypo-qwen2.5-coder-1.5b-dpo-v3
Use this script to compare the four model variants on an arbitrary prompt
without running the full 150-prompt characterization pipeline.
Important:
- This script is a smoke test for *prompt-level* behavior.
- It mirrors the evaluation prompt rendering path, but it does NOT replay the
original batch context from the published eval artifacts.
- For exact reproduction of a row from `samples.jsonl`, use
`examples/reproduce_eval_row.py` instead.
Example:
python examples/reproduce_v3.py --prompt "Write a function that returns the nth Fibonacci number."
"""
from __future__ import annotations
import argparse
import textwrap
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer
BASE_ID = "Qwen/Qwen2.5-Coder-1.5B-Instruct"
SFT_ID = "joshuasundance/mypo-qwen2.5-coder-1.5b-sft"
DPO_V2_ID = "joshuasundance/mypo-qwen2.5-coder-1.5b-dpo-v2"
DPO_V3_ID = "joshuasundance/mypo-qwen2.5-coder-1.5b-dpo-v3"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Compare MyPO generations side-by-side")
parser.add_argument(
"--prompt",
default="Write a function that returns the nth Fibonacci number.",
help="User prompt to run through all four models.",
)
parser.add_argument("--max-new-tokens", type=int, default=384)
return parser.parse_args()
def build_inputs(tokenizer: AutoTokenizer, prompt: str, device: torch.device):
# Mirror mypo_generate.py as closely as possible: user-only chat template,
# then tokenize the rendered prompt with left padding and truncation.
messages = [{"role": "user", "content": prompt}]
rendered = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = tokenizer(
[rendered],
return_tensors="pt",
padding=True,
truncation=True,
max_length=2048,
)
return inputs.to(device)
def generate(model: AutoModelForCausalLM, tokenizer: AutoTokenizer, prompt: str, max_new_tokens: int) -> str:
inputs = build_inputs(tokenizer, prompt, model.device)
output = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=tokenizer.pad_token_id,
)
prompt_length = inputs["input_ids"].shape[-1]
return tokenizer.decode(output[0][prompt_length:], skip_special_tokens=True).strip()
def banner(title: str) -> str:
return f"\n{'=' * 20} {title} {'=' * 20}\n"
def main() -> int:
args = parse_args()
tokenizer = AutoTokenizer.from_pretrained(BASE_ID)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"
print("Loading base model...", flush=True)
base = AutoModelForCausalLM.from_pretrained(
BASE_ID,
dtype=torch.bfloat16,
device_map="auto",
attn_implementation="sdpa",
)
base.eval()
print("Loading SFT and DPO-v2 adapters on the shared base...", flush=True)
peft_model = PeftModel.from_pretrained(base, SFT_ID, adapter_name="sft")
peft_model.load_adapter(DPO_V2_ID, adapter_name="dpo_v2")
peft_model.eval()
print("Loading merged DPO-v3 model...", flush=True)
dpo_v3 = AutoModelForCausalLM.from_pretrained(
DPO_V3_ID,
dtype=torch.bfloat16,
device_map="auto",
attn_implementation="sdpa",
)
dpo_v3.eval()
print(banner("PROMPT"))
print(textwrap.fill(args.prompt, width=100))
print(banner("BASE"))
print(generate(base, tokenizer, args.prompt, args.max_new_tokens))
peft_model.set_adapter("sft")
print(banner("SFT"))
print(generate(peft_model, tokenizer, args.prompt, args.max_new_tokens))
peft_model.set_adapter("dpo_v2")
print(banner("DPO-V2"))
print(generate(peft_model, tokenizer, args.prompt, args.max_new_tokens))
print(banner("DPO-V3"))
print(generate(dpo_v3, tokenizer, args.prompt, args.max_new_tokens))
return 0
if __name__ == "__main__":
raise SystemExit(main())