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