"""Step 6 — hold out eval prompts, so the eval set is not scored on its own training data. python training/06_split_eval.py Master prompt 5.4 wants 200 held-out prompts: 100 GF across the moods and 100 Normal. P1 ships 66 hand-written ones carrying the hard checks; this tops them up from the restyled data and writes the rest to brain/evals/prompts/held_out_*.yaml. Held out means held out: every id written here is recorded in `data/eval-ids.txt`, and 05_embed_load.py refuses to load those rows into the style bank. Scoring a model on examples it was shown measures memorisation. """ import argparse import random import sys from collections import defaultdict from pathlib import Path import yaml sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from training.pipeline import buckets, restyle, review # noqa: E402 ROOT = Path(__file__).parent STYLED = ROOT / "data" / "styled" REVIEW = ROOT / "review" EVAL_IDS = ROOT / "data" / "eval-ids.txt" PROMPTS = Path(__file__).resolve().parents[1] / "brain" / "evals" / "prompts" TARGET_PER_MODE = 100 #: P1's hand-written set already covers this many, so only the remainder is held out. HAND_WRITTEN = {"gf": 35, "normal": 31} def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--seed", type=int, default=7) args = parser.parse_args() decision = ( review.parse_dir(REVIEW) if REVIEW.exists() else review.ReviewDecision(frozenset(), 0) ) rng = random.Random(args.seed) by_mode: dict[str, list[dict]] = defaultdict(list) for bucket in buckets.BUCKETS: for row in restyle.read_jsonl(STYLED / f"{bucket.name}.jsonl"): if decision.keeps(row["id"]): by_mode[row["mode"]].append(row) if not by_mode: print("Nothing restyled yet. Run 03_restyle.py first.") return 1 held: list[str] = [] for mode, rows in by_mode.items(): wanted = max(0, TARGET_PER_MODE - HAND_WRITTEN.get(mode, 0)) if len(rows) < wanted: print(f"{mode}: only {len(rows)} rows available, wanted {wanted}.") wanted = len(rows) # Stratified by mood, so the held-out set is not accidentally all `neutral`. by_mood: dict[str, list[dict]] = defaultdict(list) for row in rows: by_mood[row.get("mood") or ""].append(row) picked: list[dict] = [] moods = sorted(by_mood) while len(picked) < wanted and any(by_mood.values()): for mood in moods: if by_mood[mood] and len(picked) < wanted: picked.append(by_mood[mood].pop(rng.randrange(len(by_mood[mood])))) cases = [ { "id": f"held-{row['id']}", "mood": row.get("mood") or "neutral", "category": "held_out", "prompt": (row.get("context") or [""])[-1], } for row in picked if (row.get("context") or [""])[-1].strip() ] path = PROMPTS / f"held_out_{mode}.yaml" path.write_text( "# Held out by training/06_split_eval.py. Never loaded into the style bank.\n" + yaml.safe_dump(cases, allow_unicode=True, sort_keys=False), encoding="utf-8", ) held += [row["id"] for row in picked] print(f"{mode}: {len(cases)} held out → {path}") EVAL_IDS.parent.mkdir(parents=True, exist_ok=True) EVAL_IDS.write_text("\n".join(sorted(held)) + "\n", encoding="utf-8") print(f"\n{len(held)} ids recorded in {EVAL_IDS}; 05_embed_load.py will skip them.") return 0 if __name__ == "__main__": raise SystemExit(main())