creative-writing-llm / src /build_pool.py
Pranav2748's picture
Add src
cbc33fe verified
Raw
History Blame Contribute Delete
6.76 kB
"""
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())