""" Build the scored generation pool that feeds E3 (multi-positive diverse DPO) and E4 (faithful DivPO), and doubles as base-model analysis data. Per prompt: N=16 samples at temperature 1.0 from the BASE policy, each carrying - text, token count, cumulative + length-normalized logprob (E4 divpo-prob) - programmatic gate result - judge quality / novelty (gate-passing stories only) - embedding, per-group deviation d_i and marginal contribution m_i Both DPO arms consume this identical artifact, which is the point: E4-vs-E3 is then a comparison of PAIR CONSTRUCTION and LOSS, with the data held fixed. """ from __future__ import annotations import argparse import json import sys import time from pathlib import Path import numpy as np ROOT = Path(__file__).resolve().parent.parent def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", default="Qwen/Qwen3-4B-Instruct-2507") ap.add_argument("--tag", default="4b") ap.add_argument("--n", type=int, default=16) ap.add_argument("--temperature", type=float, default=1.0) ap.add_argument("--top-p", type=float, default=1.0) ap.add_argument("--limit", type=int, default=None, help="prompt subset, for smoke") ap.add_argument("--split", default="train") ap.add_argument("--seed", type=int, default=1234) ap.add_argument("--gpu-mem", type=float, default=0.85) args = ap.parse_args() from transformers import AutoTokenizer import gates import logbook from data import load_prompts from diversity import (l2_normalize, logdet_volume, marginal_contributions, pairwise_deviation) from generate import build_llm, generate from judge import build_judge out_dir = ROOT / "outputs" / f"pool_{args.tag}" out_dir.mkdir(parents=True, exist_ok=True) pool_path = out_dir / f"pool_{args.split}.jsonl" prompts = load_prompts(args.split, ROOT / "data") if args.limit: prompts = prompts[: args.limit] print(f"[pool] {len(prompts)} prompts x N={args.n} = {len(prompts)*args.n} stories") # ---- 1. generate ----------------------------------------------------- t0 = time.time() tok = AutoTokenizer.from_pretrained(args.model) llm = build_llm(args.model, gpu_mem_util=args.gpu_mem, seed=args.seed) gens = generate(llm, tok, prompts, n=args.n, temperature=args.temperature, top_p=args.top_p, seed=args.seed) t_gen = time.time() - t0 ntok = sum(g.n_tokens for g in gens) print(f"[gen] {len(gens)} stories, {ntok} tok in {t_gen/60:.1f} min " f"({ntok/t_gen:.0f} tok/s)") # free the GPU before loading the embedder del llm import gc, torch gc.collect(); torch.cuda.empty_cache() # ---- 2. gates -------------------------------------------------------- grs = [gates.check(g.text, finish_reason=g.finish_reason) for g in gens] n_pass = sum(r.passed for r in grs) print(f"[gates] pass {n_pass}/{len(grs)} ({100*n_pass/len(grs):.1f}%)") # ---- 3. judge (gate-passers only) ------------------------------------ judge = build_judge(cache_path=str(ROOT / "cache" / "judge.sqlite"), concurrency=24) idx = [i for i in range(len(gens)) if grs[i].passed] t1 = time.time() scores = judge.score_many_sync([(gens[i].prompt, gens[i].text) for i in idx]) print(f"[judge] {len(idx)} scored in {(time.time()-t1)/60:.1f} min | " f"health={judge.health()}") judge.assert_healthy() quality = np.zeros(len(gens)); novelty = np.zeros(len(gens)) for i, s in zip(idx, scores): quality[i] = s.quality; novelty[i] = s.novelty # ---- 4. embeddings + per-group diversity ----------------------------- from sentence_transformers import SentenceTransformer enc = SentenceTransformer("BAAI/bge-base-en-v1.5", device="cuda") t2 = time.time() E = enc.encode([g.text for g in gens], normalize_embeddings=True, batch_size=64, show_progress_bar=False, convert_to_numpy=True) E = l2_normalize(np.asarray(E, dtype=np.float64)) print(f"[embed] {E.shape} in {time.time()-t2:.0f}s") by_prompt: dict[str, list[int]] = {} for i, g in enumerate(gens): by_prompt.setdefault(g.prompt_id, []).append(i) dev = np.zeros(len(gens)); marg = np.zeros(len(gens)) group_logdet: dict[str, float] = {} for pid, ids in by_prompt.items(): sub = E[ids] d = pairwise_deviation(sub); m = marginal_contributions(sub) for k, i in enumerate(ids): dev[i] = d[k]; marg[i] = m[k] group_logdet[pid] = logdet_volume(sub) # ---- 5. write -------------------------------------------------------- with open(pool_path, "w") as f: for i, g in enumerate(gens): f.write(json.dumps({ **g.as_dict(), "gate_passed": bool(grs[i].passed), "gate_reasons": grs[i].reasons, "ends_cleanly": grs[i].completeness, "n_words": grs[i].n_words, "quality": float(quality[i]), "novelty": float(novelty[i]), "deviation": float(dev[i]), "marginal": float(marg[i]), "group_logdet": float(group_logdet[g.prompt_id]), }) + "\n") np.save(out_dir / f"emb_{args.split}.npy", E.astype(np.float32)) # ---- 6. summary ------------------------------------------------------ q_pass = quality[[i for i in idx]] summary = { "tag": args.tag, "model": args.model, "split": args.split, "n_prompts": len(prompts), "n_per_prompt": args.n, "n_stories": len(gens), "temperature": args.temperature, "seed": args.seed, "gate_pass_rate": float(n_pass / len(grs)), "ends_cleanly_rate": float(np.mean([r.completeness for r in grs])), "median_words": float(np.median([r.n_words for r in grs])), "quality_mean": float(q_pass.mean()) if len(q_pass) else 0.0, "quality_sd": float(q_pass.std()) if len(q_pass) else 0.0, "novelty_mean": float(novelty[idx].mean()) if len(idx) else 0.0, "deviation_mean": float(dev.mean()), "logdet_mean": float(np.mean(list(group_logdet.values()))), "gen_minutes": t_gen / 60, "judge_cost": judge.cost_estimate(0.140, 0.280), "judge_health": judge.health(), } json.dump(summary, open(out_dir / f"summary_{args.split}.json", "w"), indent=2) print("\n" + json.dumps(summary, indent=1)) logbook.note(f"pool built: {args.tag}/{args.split}", f"```json\n{json.dumps(summary, indent=1)}\n```") logbook.checkpoint(f"pool_{args.tag}_{args.split}") return 0 if __name__ == "__main__": sys.exit(main())