Download jev_toy/eval.py from azharmo/build-jev-from-scratch: direct link, hf CLI and curl.
- Browser
- Download file 4.71 kB
-
https://huggingface.co/azharmo/build-jev-from-scratch/resolve/f48dabc6e187bfc9c7b5a9c2699c5320667730d4/jev_toy/eval.py
- Command line
-
hf download hf://azharmo/build-jev-from-scratch@f48dabc6e187bfc9c7b5a9c2699c5320667730d4/jev_toy/eval.py
-
curl -L -o eval.py https://huggingface.co/azharmo/build-jev-from-scratch/resolve/f48dabc6e187bfc9c7b5a9c2699c5320667730d4/jev_toy/eval.py
4.71 kB
| """ | |
| jev_toy/eval.py | |
| Measure what matters for a System One decision model: | |
| - accuracy (does the argmax/label match) | |
| - expected calibration error (ECE) and reliability | |
| - Brier score (proper scoring rule) | |
| Applied to noul and choice heads. Also does post-hoc temperature scaling | |
| (our implementation of the "calibrated" property; tuning T on a held-out set). | |
| Usage: | |
| python -m jev_toy.eval --ckpt checkpoints/model.pt | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import math | |
| import torch | |
| import torch.nn.functional as F | |
| from jev_toy.data import build_examples, tokenize | |
| from jev_toy.model import SystemOneConfig, SystemOneModel | |
| from jev_toy.train import featurize | |
| def predict(model, eval_examples, vocab, cfg, device): | |
| """Return per-row predictions for noul and choice rows.""" | |
| model.eval() | |
| model = model.to(device) | |
| n_rows, c_rows = [], [] | |
| for bidx in range(0, len(eval_examples), 64): | |
| idx = list(range(bidx, min(bidx + 64, len(eval_examples)))) | |
| (ids_s, mask_s, ids_q, mask_q, types, targets) = featurize(idx, eval_examples, vocab, cfg, device) | |
| h = model.encode_state(ids_s, mask_s) | |
| logits, _ = model.answer(h, ids_q, mask_q, types) | |
| for k, t in enumerate(types): | |
| if t == "noul": | |
| n_rows.append({"p": torch.sigmoid(logits["noul"][k]).item(), "y": int(targets[k])}) | |
| elif t == "choice": | |
| probs = F.softmax(logits["choice"][k], dim=-1).cpu() | |
| c_rows.append({"p": probs, "y": int(targets[k])}) | |
| return n_rows, c_rows | |
| def brier(p, y): | |
| return (p - y) ** 2 | |
| def ece(probs, ys, n_bins=10): | |
| """Expected Calibration Error over probability bins.""" | |
| bins = [0.0] * n_bins | |
| conf = [0.0] * n_bins | |
| cnt = [0] * n_bins | |
| for p, y in zip(probs, ys): | |
| b = min(int(p * n_bins), n_bins - 1) | |
| cnt[b] += 1 | |
| conf[b] += p | |
| bins[b] += 1.0 if p >= 0.5 and y == 1 or p < 0.5 and y == 0 else 0.0 | |
| tot = sum(cnt) | |
| if tot == 0: | |
| return 0.0 | |
| e = 0.0 | |
| for i in range(n_bins): | |
| if cnt[i]: | |
| acc = bins[i] / cnt[i] | |
| e += cnt[i] / tot * abs(acc - conf[i] / cnt[i]) | |
| return e | |
| def noul_metrics(n_rows, temp=1.0): | |
| probs = [r["p"] for r in n_rows] | |
| ys = [r["y"] for r in n_rows] | |
| if temp != 1.0: | |
| probs = [1.0 / (1.0 + math.exp(-(math.log(p / (1 - p + 1e-9))) / temp)) for p in probs] | |
| acc = sum(1 for p, y in zip(probs, ys) if (p >= 0.5) == (y == 1)) / len(n_rows) | |
| brier = sum((p - y) ** 2 for p, y in zip(probs, ys)) / len(n_rows) | |
| return {"n": len(n_rows), "acc": acc, "brier": brier, "ece": ece(probs, ys)} | |
| def choice_metrics(c_rows, temp=1.0): | |
| acc = 0 | |
| brier_total = 0.0 | |
| for r in c_rows: | |
| logit = torch.log(r["p"] + 1e-9) / temp | |
| probs = F.softmax(logit, dim=-1) | |
| pred = int(probs.argmax()) | |
| acc += (pred == r["y"]) | |
| y1h = torch.zeros_like(probs) | |
| y1h[r["y"]] = 1.0 | |
| brier_total += ((probs - y1h) ** 2).sum().item() | |
| n = len(c_rows) | |
| return {"n": n, "acc": acc / n, "brier": brier_total / n} | |
| def temperature_scan(n_rows, c_rows): | |
| """Pick the temperature that minimises ECE on a held-out split.""" | |
| half_n = len(n_rows) // 2 | |
| half_c = len(c_rows) // 2 | |
| best = {"nce": (1.0, 9e9), "cce": (1.0, 9e9)} | |
| for t in [0.4, 0.6, 0.8, 1.0, 1.2, 1.5, 2.0, 3.0]: | |
| mn = noul_metrics(n_rows[:half_n], t) | |
| if mn["ece"] < best["nce"][1]: | |
| best["nce"] = (t, mn["ece"]) | |
| return best | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--ckpt", default="checkpoints/model.pt") | |
| args = ap.parse_args() | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| ck = torch.load(args.ckpt, map_location=device) | |
| cfg = SystemOneConfig(**ck["config"]) | |
| vocab = type("V", (), {"stoi": ck["vocab"], "oov": ck["vocab"].get("__oov__", 1)})() | |
| model = SystemOneModel(cfg).to(device) | |
| model.load_state_dict(ck["state_dict"]) | |
| from datasets import load_dataset | |
| ag = load_dataset("fancyzhx/ag_news") | |
| bq = load_dataset("google/boolq") | |
| st = load_dataset("stanfordnlp/sst2") | |
| # small eval splits | |
| subsample = {"agnews": 600, "boolq": 600, "sst2": 600} | |
| eval_examples, _ = build_examples({"agnews": ag, "boolq": bq, "sst2": st}, subsample) | |
| n_rows, c_rows = predict(model, eval_examples, vocab, cfg, device) | |
| print("--- noul (BoolQ+SST2) ---", noul_metrics(n_rows)) | |
| print("--- choice (AG News) ---", choice_metrics(c_rows)) | |
| ts = temperature_scan(n_rows, c_rows) | |
| print("best noul temperature:", ts["nce"][0], "ECE", round(ts["nce"][1], 4)) | |
| if __name__ == "__main__": | |
| main() |