summaryrefslogtreecommitdiff
path: root/src/eval.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/eval.py')
-rw-r--r--src/eval.py151
1 files changed, 151 insertions, 0 deletions
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()