From 898a0570dfe619ec2bcf330b07dce9b23cb63d54 Mon Sep 17 00:00:00 2001 From: Void Agent Date: Fri, 14 Aug 2026 13:11:02 +0100 Subject: implement Experiment 1: data pipeline, tied-RNN w/ ACT halting, transformer baseline, train/eval/plot, 16 tests --- src/eval.py | 151 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 151 insertions(+) create mode 100644 src/eval.py (limited to 'src/eval.py') diff --git a/src/eval.py b/src/eval.py new file mode 100644 index 0000000..9d55437 --- /dev/null +++ b/src/eval.py @@ -0,0 +1,151 @@ +"""Post-run analysis: final metrics, [101,200] probe with sieve diagnostic, grokking signature. + +Interpretation codes are locked in design/preregistration.md — this module only MEASURES +and classifies against those definitions (O1-O4, H1-H4, P1-P4). +""" +import argparse +import csv +import json +import os + +import numpy as np +import torch + +from src.config import Config +from src.data import build_examples, decode_tokens, encode_int, get_splits, next_prime, sieve_primes +from src.model_api import build_model, greedy_decode +from src.train import evaluate + +FLAGGED_COMPOSITES = {121, 143, 169, 187} # need divisors 11, 13 — beyond the {2,3,5,7} sieve + + +def probe_report(model, cfg: Config, lo: int = 101, hi: int = 200) -> dict: + primes = sieve_primes(hi + 200) + correct = 0 + errors = [] + easy_misses = 0 + for n in range(lo, hi + 1): + x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0) + gen = greedy_decode(model, x, cfg)[0].tolist() + pred = decode_tokens(gen, cfg) + target = next_prime(n, primes) + if pred == target: + correct += 1 + else: + errors.append({"n": n, "target": target, "pred": pred}) + if n % 2 == 0 or n % 5 == 0: + easy_misses += 1 + total = hi - lo + 1 + acc = correct / total + flagged = [e for e in errors if e["n"] in FLAGGED_COMPOSITES] + # P-code classification (see preregistration.md) + if acc >= 0.85: + code = "P3" + elif easy_misses >= 5: + code = "P4" + elif errors and all(e["n"] in FLAGGED_COMPOSITES for e in errors) and len(errors) <= 6: + code = "P1" + else: + code = "P2" + return { + "code": code, "acc": acc, "correct": correct, "total": total, + "errors": errors, "flagged_errors": flagged, "easy_misses": easy_misses, + } + + +def grokking_signature(metrics_path: str) -> dict: + """Classify the training curve against preregistered codes O1-O4.""" + with open(metrics_path) as fh: + rows = list(csv.DictReader(fh)) + if not rows: + return {"code": "O4", "note": "no eval rows"} + train = [float(r["train_em"]) for r in rows] + val = [float(r["val_em"]) for r in rows] + n = len(rows) + saturated = any(all(t >= 0.95 for t in train[i:i + 10]) for i in range(n - 9)) if n >= 10 else False + hi = next((i for i, v in enumerate(val) if v >= 0.9), None) + trans = None + if hi is not None: + lo_cands = [i for i in range(hi) if val[i] <= 0.2] + if lo_cands: + trans = hi - max(lo_cands) + if saturated and hi is not None and trans is not None and trans <= 5: + code = "O1" + elif max(train) >= 0.95 and hi is not None and (trans is None or trans > 5): + code = "O3" + elif max(train) >= 0.95 and hi is None: + code = "O2" + else: + code = "O4" + return { + "code": code, "train_saturated": saturated, "val_hi_eval_idx": hi, + "transition_width_evals": trans, "n_evals": n, + "final_train_em": train[-1], "final_val_em": val[-1], + } + + +@torch.no_grad() +def halting_report(model, cfg: Config) -> dict: + """RNN halting structure: mean steps + correlation with gap-to-next-prime (H1-H4).""" + primes = sieve_primes(300) + gaps, steps = [], [] + for n in range(cfg.range_start, cfg.range_end + 1): + x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0) + h = model._initial_state(x) + _, s = model._run_cell(h) + gaps.append(next_prime(n, primes) - n) + steps.append(float(s.mean())) + mean = float(np.mean(steps)) + rho = float(np.corrcoef(gaps, steps)[0, 1]) if len(set(gaps)) > 1 else 0.0 + if mean <= cfg.min_steps + 0.5: + code = "H1" + elif mean >= cfg.max_steps - 0.5: + code = "H2" + elif abs(rho) >= 0.3: + code = "H3" + else: + code = "H4" + return {"code": code, "mean_steps": mean, "corr_gap": rho, "min_steps": cfg.min_steps, "max_steps": cfg.max_steps} + + +def main() -> None: + ap = argparse.ArgumentParser(description="prime-grokking eval") + ap.add_argument("model") + ap.add_argument("seed") + ap.add_argument("--ckpt", default="best.pt") + a = ap.parse_args() + out_dir = os.path.join("runs", a.model, f"seed{a.seed}") + cfg = Config.load(os.path.join(out_dir, "config.json")) + model = build_model(cfg) + model.load_state_dict(torch.load(os.path.join(out_dir, a.ckpt), map_location="cpu")) + model.eval() + + _, val_in = get_splits(cfg) + val_ex = build_examples(val_in, cfg) + tok, em, per = evaluate(model, val_ex, cfg) + + probe = probe_report(model, cfg) + sig = grokking_signature(os.path.join(out_dir, "metrics.csv")) + hlt = halting_report(model, cfg) if cfg.model == "rnn" else None + + results = { + "model": cfg.model, "seed": cfg.seed, "ckpt": a.ckpt, + "params": model.param_count(), + "val_token_acc": tok, "val_exact_match": em, + "per_example": [{"n": "".join(map(str, x)), "target": t, "pred": p, "ok": ok} + for x, t, p, ok in per], + "probe": probe, + "signature": sig, + "halting": hlt, + } + with open(os.path.join(out_dir, "results.json"), "w") as fh: + json.dump(results, fh, indent=2) + print(json.dumps({"val_token_acc": round(tok, 4), "val_em": round(em, 4), + "probe": probe["code"], "probe_acc": round(probe["acc"], 3), + "flagged_errors": [e["n"] for e in probe["flagged_errors"]], + "signature": sig["code"], "halting": hlt["code"] if hlt else None, + "results": os.path.join(out_dir, "results.json")}, indent=2)) + + +if __name__ == "__main__": + main() -- cgit v1.2.3