"""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()