prova-core / training /06_split_eval.py
chintakapp's picture
core 2.0
210ef30
Raw History Blame Contribute Delete
3.7 kB
"""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())