""" Checkpoint trajectory study: generate the SAME prompts from every checkpoint of an arm, so the evolution of the actual stories is visible -- not just metrics. Aggregate numbers can report "effective rank 2.0" without conveying that six of sixteen ships are named *Aethel*. This dumps the stories in a readable form at each training step alongside the metrics, so the qualitative change can be read directly and cross-checked against the quantitative one. vLLM loads the base model ONCE and hot-swaps LoRA adapters per checkpoint, so the whole sweep costs one model load rather than one per checkpoint. Outputs (under outputs/ckpt_study//): stories.md human-readable: every prompt, every checkpoint, side by side metrics.csv per-checkpoint quantitative trajectory raw.json everything, for re-analysis ../logs/figures/_trajectory.png """ from __future__ import annotations import argparse import csv import json import re import sys from pathlib import Path import numpy as np ROOT = Path(__file__).resolve().parent.parent def find_checkpoints(arm: str) -> list[tuple[int, str | None]]: """[(step, adapter_path_or_None)] ascending; step 0 = base model.""" d = ROOT / "outputs" / arm out: list[tuple[int, str | None]] = [(0, None)] if d.exists(): for p in d.glob("checkpoint-*"): m = re.search(r"checkpoint-(\d+)", p.name) if m and (p / "adapter_model.safetensors").exists(): out.append((int(m.group(1)), str(p))) f = d / "final" if (f / "adapter_model.safetensors").exists(): steps = [s for s, _ in out] out.append((max(steps) + 1 if steps else 1, str(f))) return sorted(out, key=lambda t: t[0]) def main(): ap = argparse.ArgumentParser() ap.add_argument("--arm", required=True) ap.add_argument("--model", default="Qwen/Qwen3-4B-Instruct-2507") ap.add_argument("--prompts", type=int, default=10) ap.add_argument("--n", type=int, default=6) ap.add_argument("--temp", type=float, default=0.9) ap.add_argument("--top-p", type=float, default=0.95) ap.add_argument("--seed", type=int, default=777) ap.add_argument("--gpu-mem", type=float, default=0.85) ap.add_argument("--judge", action="store_true", default=True) args = ap.parse_args() from transformers import AutoTokenizer from vllm import SamplingParams from vllm.lora.request import LoRARequest import gates import logbook from data import load_prompts from diversity import effective_rank, l2_normalize, logdet_volume, pairwise_deviation from generate import build_llm, render_chat from judge import build_judge from qualitative import analyze_group, first_sentence, last_sentence ckpts = find_checkpoints(args.arm) if len(ckpts) < 2: print(f"only {len(ckpts)} checkpoint(s) for {args.arm}; nothing to compare") return 2 print(f"[{args.arm}] checkpoints: {[s for s, _ in ckpts]}") prompts = load_prompts("eval", ROOT / "data")[: args.prompts] tok = AutoTokenizer.from_pretrained(args.model) llm = build_llm(args.model, gpu_mem_util=args.gpu_mem, seed=args.seed, enable_lora=True) rendered = [render_chat(tok, p["prompt"]) for p in prompts] sp = SamplingParams(n=args.n, temperature=args.temp, top_p=args.top_p, max_tokens=1024, seed=args.seed, skip_special_tokens=True) from sentence_transformers import SentenceTransformer enc = None judge = build_judge(cache_path=str(ROOT / "cache" / "judge.sqlite"), concurrency=24) \ if args.judge else None all_data, rowsum = {}, [] for step, path in ckpts: kw = {} if path: kw["lora_request"] = LoRARequest(f"{args.arm}-{step}", max(step, 1), path) outs = llm.generate(rendered, sp, **kw) per = {} for p, o in zip(prompts, outs): texts = [x.text.strip() for x in o.outputs] frs = [x.finish_reason or "" for x in o.outputs] per[p["id"]] = {"prompt": p["prompt"], "texts": texts, "gates": [gates.check(t, finish_reason=f).as_dict() for t, f in zip(texts, frs)]} all_data[step] = per print(f" step {step:>4}: generated {sum(len(v['texts']) for v in per.values())} stories", flush=True) del llm import gc, torch gc.collect(); torch.cuda.empty_cache() enc = SentenceTransformer("BAAI/bge-base-en-v1.5", device="cuda") for step, per in all_data.items(): dev, ld, er, q = [], [], [], [] pooled_texts = [] for pid, v in per.items(): E = l2_normalize(np.asarray(enc.encode( v["texts"], normalize_embeddings=True, show_progress_bar=False, convert_to_numpy=True), dtype=np.float64)) v["eff_rank"] = float(effective_rank(E)) v["deviation"] = float(pairwise_deviation(E).mean()) v["logdet"] = float(logdet_volume(E)) v["qual"] = analyze_group(v["texts"]) dev.append(v["deviation"]); ld.append(v["logdet"]); er.append(v["eff_rank"]) pooled_texts += [(v["prompt"], t, g["passed"]) for t, g in zip(v["texts"], v["gates"])] if judge: idx = [i for i, (_, _, ok) in enumerate(pooled_texts) if ok] sc = judge.score_many_sync([(pooled_texts[i][0], pooled_texts[i][1]) for i in idx]) q = [s.quality for s in sc] gp = float(np.mean([g["passed"] for v in per.values() for g in v["gates"]])) ec = float(np.mean([g["completeness"] for v in per.values() for g in v["gates"]])) wd = float(np.mean([g["n_words"] for v in per.values() for g in v["gates"]])) rowsum.append({ "step": step, "quality": float(np.mean(q)) if q else 0.0, "eff_rank": float(np.mean(er)), "deviation": float(np.mean(dev)), "logdet": float(np.mean(ld)), "gate_pass": gp, "ends_cleanly": ec, "words": wd, "opens_with_The": float(np.mean([v["qual"]["opens_with_The"] / v["qual"]["n"] for v in per.values()])), "distinct_openers": float(np.mean([v["qual"]["distinct_first_5_words"] / v["qual"]["n"] for v in per.values()])), "registers": float(np.mean([v["qual"]["registers_present"] for v in per.values()])), }) print(f" step {step:>4}: q={rowsum[-1]['quality']:.2f} " f"eff_rank={rowsum[-1]['eff_rank']:.3f} dev={rowsum[-1]['deviation']:.4f} " f"words={wd:.0f}", flush=True) out = ROOT / "outputs" / "ckpt_study" / args.arm out.mkdir(parents=True, exist_ok=True) with open(out / "metrics.csv", "w", newline="") as f: w = csv.DictWriter(f, fieldnames=list(rowsum[0].keys())) w.writeheader(); w.writerows(rowsum) json.dump(all_data, open(out / "raw.json", "w"), indent=1) # ---- human-readable side-by-side -------------------------------------- steps = [s for s, _ in ckpts] md = [f"# {args.arm} — story trajectory across checkpoints\n", f"{len(prompts)} eval prompts x {args.n} samples, T={args.temp}, " f"top_p={args.top_p}, seed={args.seed} (fixed across checkpoints).\n", "Step 0 = base model.\n"] for pid in list(all_data[steps[0]]): md.append(f"\n## {pid}\n\n> {all_data[steps[0]][pid]['prompt']}\n") for s in steps: v = all_data[s][pid] md.append(f"\n### step {s} — eff_rank {v['eff_rank']:.2f}, " f"dev {v['deviation']:.3f}\n") md.append("\n**openings**\n") for i, t in enumerate(v["texts"]): md.append(f"{i+1}. {first_sentence(t, 150)}\n") md.append("\n**closings**\n") for i, t in enumerate(v["texts"]): md.append(f"{i+1}. …{last_sentence(t, 110)}\n") (out / "stories.md").write_text("".join(md)) # ---- figure ----------------------------------------------------------- import matplotlib; matplotlib.use("Agg") import matplotlib.pyplot as plt x = [r["step"] for r in rowsum] fig, ax = plt.subplots(1, 4, figsize=(19, 4.2)) # Anchor the axes that have a meaningful absolute scale. Auto-scaling a # metric that moved 1.60->1.69 against a ceiling of N renders a dramatic # line for a flat result, which is exactly the misreading to avoid. for a, (k, t, c) in zip(ax, [("quality", "Judge quality (0-10)", "#2980b9"), ("eff_rank", "Effective rank (1 = collapsed, %d = max)" % args.n, "#8e44ad"), ("deviation", "Mean pairwise deviation", "#16a085"), ("words", "Story length (words)", "#7f8c8d")]): vals = [r[k] for r in rowsum] a.plot(x, vals, "o-", color=c, lw=2) if k == "eff_rank": a.set_ylim(1.0, args.n) # full meaningful range a.axhline(1.0, ls=":", c="crimson", lw=1) a.text(x[0], 1.05, "total collapse", fontsize=7, color="crimson") elif k == "quality": a.set_ylim(0, 10) elif k == "deviation": a.set_ylim(0, max(0.5, max(vals) * 1.3)) a.set_title(t, fontsize=10); a.set_xlabel("training step"); a.grid(alpha=.3) fig.suptitle(f"{args.arm}: what happens to the stories during training", fontsize=13) plt.tight_layout() figp = logbook.FIGS / f"{args.arm}_trajectory.png" plt.savefig(figp, dpi=140); plt.close() print(f"\nstories -> {out/'stories.md'}") print(f"metrics -> {out/'metrics.csv'}") print(f"figure -> {figp}") if judge: print("judge:", judge.health(), judge.cost_estimate(0.140, 0.280)) return 0 if __name__ == "__main__": sys.exit(main())