summaryrefslogtreecommitdiff
path: root/src/eval.py
blob: 9d554374d3175fbbe9006ff87f1dcde35ad4f54e (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
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()