Download eval/evaluate.py from nur-dev/pns-bind-25m: direct link, hf CLI and curl.
- Browser
- Download file 13.9 kB
-
https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/evaluate.py
- Command line
-
hf download hf://nur-dev/pns-bind-25m/eval/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/evaluate.py
13.9 kB
| #!/usr/bin/env python3 | |
| """Streaming evaluation with causal interventions. | |
| Warm full-lifetime protocol: PNSR state is carried across the WHOLE lifetime | |
| (detach != reset; the parent project's stale-gate lesson). Value-match | |
| scoring: a pointer answer is correct iff the selected record holds the gold | |
| value bytes. | |
| Interventions (eval-time, preregistered): | |
| none | sigma_zero | reset32 | reset64 | jobs_zero | jself_zero | swap | |
| --K overrides deliberation depth on a trained model (parameter-shared). | |
| """ | |
| import argparse | |
| import json | |
| import sys | |
| import time | |
| from collections import defaultdict | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src")) | |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) | |
| from pns.common import atomic_write_json, eval_root, shards_root # noqa: E402 | |
| from pns.checkpoint import load_model as _load # noqa: E402 | |
| from pns.train.loader import iter_eval_batches, WindowSampler # noqa: E402 | |
| from pns.data.view import Shard, shard_paths # noqa: E402 | |
| from pns.world.schema import Fam, Mode # noqa: E402 | |
| BUCKETS = [(1, 8), (9, 24), (25, 48), (49, 128), (129, 512), (513, 4096)] | |
| def bucket(d): | |
| for i, (lo, hi) in enumerate(BUCKETS): | |
| if lo <= d <= hi: | |
| return i | |
| return -1 | |
| def to_gpu(b, dev): | |
| g = {} | |
| for k, v in b.items(): | |
| if k in ("rec_val_hash", "gold_val_hash", "seed"): | |
| g[k] = torch.from_numpy(v.astype(np.int64) if v.dtype != np.uint64 | |
| else v.view(np.int64).copy()).to(dev) | |
| else: | |
| g[k] = torch.from_numpy(np.ascontiguousarray(v.astype(np.int64))).to(dev) | |
| return g | |
| def load_model(run, dev, ckpt=None): | |
| """Load a published checkpoint (safetensors + config.json).""" | |
| m, kind, _ = _load(run, dev) | |
| return m, kind | |
| def eval_pnsr(model, args, dev): | |
| rows = [] | |
| d = model.cfg.d | |
| interv = args.intervention | |
| t_events, n_events = 0.0, 0 | |
| for b in iter_eval_batches(args.split, shards_root(), args.batch, args.limit): | |
| g = to_gpu(b, dev) | |
| B, L = g["etype"].shape | |
| state = model.initial_state(B, dev) | |
| if interv == "sigma_zero": | |
| state = torch.zeros_like(state) | |
| swap_at = L // 2 | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| h_static = model.recenc.static(g["rec_val_toks"], g["rec_key_toks"], | |
| g["rec_store"], g["rec_kind"], | |
| g["rec_key"], g["rec_ent"]) | |
| for t in range(L): | |
| if interv in ("reset32", "reset64") and t > 0 and \ | |
| t % int(interv[-2:]) == 0: | |
| state = model.initial_state(B, dev) | |
| if interv == "swap" and t == swap_at: | |
| state = state.roll(1, dims=0) # donor = previous lifetime | |
| live = g["live"][:, t].clone() | |
| if interv == "jobs_zero": | |
| live[:, :64] = -1 | |
| if interv == "jself_zero": | |
| live[:, 64:88] = -1 | |
| mask = live >= 0 | |
| rowsg = live.clamp(min=0) | |
| bank = torch.gather(h_static, 1, rowsg.unsqueeze(-1).expand(-1, -1, d)) | |
| births = torch.gather(g["rec_birth"], 1, rowsg) | |
| bank = model.recenc.finalize(bank, (t - births).clamp(min=0)) | |
| bank = bank * mask.unsqueeze(-1) | |
| t0 = time.perf_counter() | |
| state, out = model.step(state, g["tok"][:, t], g["etype"][:, t], | |
| g["dt"][:, t], bank, mask, | |
| K=args.K, | |
| freeze_writes=(interv == "sigma_zero")) | |
| torch.cuda.synchronize() | |
| t_events += time.perf_counter() - t0 | |
| n_events += B | |
| score_positions(rows, out, g, t, live, b, swap_at) | |
| return rows, dict(per_event_ms=round(1000 * t_events / max(1, n_events / args.batch), 2), | |
| peak_mem_gb=round(torch.cuda.max_memory_allocated() / 2**30, 2)) | |
| def score_positions(rows, out, g, t, live, b, swap_at): | |
| mg = g["mode_gold"][:, t] | |
| interesting = mg > 0 | |
| if not interesting.any(): | |
| return | |
| idx = torch.nonzero(interesting).flatten() | |
| mode_pred = out["mode"].argmax(-1) | |
| for i in idx.tolist(): | |
| gold_mode = int(mg[i]) | |
| fam = int(g["family"][i, t]) | |
| ok = False | |
| if gold_mode == int(Mode.ANSWER_POINTER): | |
| slot = int(out["ptr"][i].argmax()) | |
| row = int(live[i, slot]) | |
| ok = row >= 0 and int(g["rec_val_hash"][i, row]) == int(g["gold_val_hash"][i, t]) | |
| elif gold_mode == int(Mode.ANSWER_ENUM): | |
| from pns.model.modules import enum_legal_mask | |
| legal = enum_legal_mask(g["enum_legal"][i:i + 1, t]) | |
| pred = int(out["enum"][i].masked_fill(~legal[0], float("-inf")).argmax()) | |
| ok = pred == int(g["enum_gold"][i, t]) | |
| elif gold_mode == int(Mode.EXTERNAL_OPERATION): | |
| okop = int(out["op"][i].argmax()) == int(g["op_gold"][i, t]) | |
| ok = okop | |
| for a in range(2): | |
| tgt = int(g["op_arg_slots"][i, t, a]) | |
| if tgt >= 0: | |
| slot = int(out["args"][i, a].argmax()) | |
| row = int(live[i, slot]) | |
| trow = int(live[i, tgt]) | |
| ok = ok and row >= 0 and trow >= 0 and \ | |
| int(g["rec_val_hash"][i, row]) == int(g["rec_val_hash"][i, trow]) | |
| rows.append(dict( | |
| fam=fam, ok=int(ok), delay=int(g["delay"][i, t]), | |
| corrected=int(g["corrected"][i, t]), reverted=int(g["reverted"][i, t]), | |
| mode_ok=int(int(mode_pred[i]) == gold_mode), | |
| lt=int(g["seed"][i]), t=t, | |
| post_swap_evidence=(-1 if int(g["delay"][i, t]) < 0 | |
| else int(t - int(g["delay"][i, t]) >= swap_at)), | |
| )) | |
| def eval_tx(model, args, dev): | |
| """Evaluate TX at every question position of the split (all positions, | |
| deterministic; reuses WindowSampler's builder on full shards).""" | |
| rows = [] | |
| W = model.cfg.window | |
| d = model.cfg.d | |
| t_fwd, n_fwd = 0.0, 0 | |
| for p in shard_paths(args.split, shards_root()): | |
| sh = Shard(p) | |
| ws = WindowSampler.__new__(WindowSampler) # reuse _build only | |
| ws.window = W | |
| n_l = sh.n_lifetimes if not args.limit else min(sh.n_lifetimes, args.limit) | |
| chunk = [] | |
| for i in range(n_l): | |
| lo, hi = int(sh.lt_off[i]), int(sh.lt_off[i + 1]) | |
| for e in np.nonzero(sh.mode_gold[lo:hi] > 0)[0]: | |
| chunk.append((i, int(e))) | |
| for k in range(0, len(chunk), args.batch): | |
| b = ws._build(sh, sorted(chunk[k:k + args.batch])) | |
| g = to_gpu(b, dev) | |
| # Exact-store lesions, applied to the live cache map before the bank | |
| # is gathered AND before scoring, so a pointer answer cannot resolve | |
| # through a lesioned row. The slot ranges match eval_pnsr. | |
| if args.intervention in ("jobs_zero", "jself_zero"): | |
| live = g["live"].clone() | |
| if args.intervention == "jobs_zero": | |
| live[:, :64] = -1 | |
| else: | |
| live[:, 64:88] = -1 | |
| g["live"] = live | |
| mask = g["live"] >= 0 | |
| rowsg = g["live"].clamp(min=0) | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| h_static = model.recenc.static(g["rec_val_toks"], g["rec_key_toks"], | |
| g["rec_store"], g["rec_kind"], | |
| g["rec_key"], g["rec_ent"]) | |
| bank = torch.gather(h_static, 1, rowsg.unsqueeze(-1).expand(-1, -1, d)) | |
| births = torch.gather(g["rec_birth"], 1, rowsg) | |
| bank = model.recenc.finalize(bank, (g["ev_idx"].unsqueeze(1) - births) | |
| .clamp(min=0)) * mask.unsqueeze(-1) | |
| t0 = time.perf_counter() | |
| out = model(g["tok"], bank, mask) | |
| torch.cuda.synchronize() | |
| t_fwd += time.perf_counter() - t0 | |
| n_fwd += g["tok"].shape[0] | |
| score_tx(rows, out, g) | |
| if args.limit and len({r["lt"] for r in rows}) >= args.limit: | |
| break | |
| return rows, dict(per_query_ms=round(1000 * t_fwd / max(n_fwd, 1), 3), | |
| peak_mem_gb=round(torch.cuda.max_memory_allocated() / 2**30, 2)) | |
| def score_tx(rows, out, g): | |
| from pns.model.modules import enum_legal_mask | |
| B = g["tok"].shape[0] | |
| mode_pred = out["mode"].argmax(-1) | |
| for i in range(B): | |
| gold_mode = int(g["mode_gold"][i]) | |
| ok = False | |
| if gold_mode == int(Mode.ANSWER_POINTER): | |
| slot = int(out["ptr"][i].argmax()) | |
| row = int(g["live"][i, slot]) | |
| ok = row >= 0 and int(g["rec_val_hash"][i, row]) == int(g["gold_val_hash"][i]) | |
| elif gold_mode == int(Mode.ANSWER_ENUM): | |
| legal = enum_legal_mask(g["enum_legal"][i:i + 1]) | |
| pred = int(out["enum"][i].masked_fill(~legal[0], float("-inf")).argmax()) | |
| ok = pred == int(g["enum_gold"][i]) | |
| elif gold_mode == int(Mode.EXTERNAL_OPERATION): | |
| ok = int(out["op"][i].argmax()) == int(g["op_gold"][i]) | |
| for a in range(2): | |
| tgt = int(g["op_arg_slots"][i, a]) | |
| if tgt >= 0: | |
| slot = int(out["args"][i, a].argmax()) | |
| row, trow = int(g["live"][i, slot]), int(g["live"][i, tgt]) | |
| ok = ok and row >= 0 and trow >= 0 and \ | |
| int(g["rec_val_hash"][i, row]) == int(g["rec_val_hash"][i, trow]) | |
| else: | |
| continue | |
| rows.append(dict(fam=int(g["family"][i]), ok=int(ok), delay=int(g["delay"][i]), | |
| corrected=int(g["corrected"][i]), reverted=int(g["reverted"][i]), | |
| mode_ok=int(int(mode_pred[i]) == gold_mode), | |
| lt=int(g["lt_seed"][i]), t=int(g["ev_idx"][i]), | |
| post_swap_evidence=-1)) | |
| MEMHARD_ENUM = {int(Fam.SEM_LATEST), int(Fam.SEM_2HOP), int(Fam.TEMPORAL_ORDER), | |
| int(Fam.DEADLINE)} | |
| def memhard(r): | |
| f = r["fam"] | |
| if f in MEMHARD_ENUM or f == int(Fam.GOAL_TOP): | |
| return True | |
| if f in (int(Fam.EXACT_DELAYED), int(Fam.EXACT_2HOP)) and r["reverted"] == 1: | |
| return True | |
| return False | |
| def aggregate(rows): | |
| agg = defaultdict(lambda: [0, 0]) | |
| def add(key, ok): | |
| agg[key][0] += ok | |
| agg[key][1] += 1 | |
| for r in rows: | |
| fam = Fam(r["fam"]).name | |
| add(f"fam/{fam}", r["ok"]) | |
| if r["fam"] in (int(Fam.EXACT_DELAYED), int(Fam.EXACT_2HOP)): | |
| add(f"fam/{fam}/rev{r['reverted']}", r["ok"]) | |
| b = bucket(r["delay"]) | |
| if b >= 0: | |
| add(f"bucket/{b}", r["ok"]) | |
| if memhard(r): | |
| add(f"memhard_bucket/{b}", r["ok"]) | |
| if memhard(r): | |
| add("memhard", r["ok"]) | |
| if r["delay"] > 48: | |
| add("memhard_beyond48", r["ok"]) | |
| if 1 <= r["delay"] <= 24: | |
| add("memhard_within24", r["ok"]) | |
| if r["post_swap_evidence"] == 0: | |
| add("memhard_preswap_evidence", r["ok"]) | |
| elif r["post_swap_evidence"] == 1: | |
| add("memhard_postswap_evidence", r["ok"]) | |
| add("mode_acc", r["mode_ok"]) | |
| add("all", r["ok"]) | |
| return {k: dict(acc=round(v[0] / v[1], 4), n=v[1]) | |
| for k, v in sorted(agg.items())} | |
| def boot_ci(rows, pred, iters=1000, seed=0): | |
| per_lt = defaultdict(lambda: [0, 0]) | |
| for r in rows: | |
| if pred(r): | |
| per_lt[r["lt"]][0] += r["ok"] | |
| per_lt[r["lt"]][1] += 1 | |
| lts = list(per_lt.values()) | |
| if not lts: | |
| return None | |
| rng = np.random.default_rng(seed) | |
| accs = [] | |
| for _ in range(iters): | |
| pick = rng.integers(0, len(lts), len(lts)) | |
| ok = sum(lts[i][0] for i in pick) | |
| n = sum(lts[i][1] for i in pick) | |
| accs.append(ok / max(n, 1)) | |
| return [round(float(np.percentile(accs, q)), 4) for q in (2.5, 97.5)] | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--run", required=True) | |
| ap.add_argument("--split", default="val") | |
| ap.add_argument("--limit", type=int, default=None) | |
| ap.add_argument("--batch", type=int, default=64) | |
| ap.add_argument("--intervention", default="none") | |
| ap.add_argument("--K", type=int, default=None) | |
| ap.add_argument("--tag", default=None) | |
| args = ap.parse_args() | |
| dev = "cuda" | |
| model, kind = load_model(args.run, dev) | |
| if kind in ("pnsr", "pnsr_k1", "rmt"): | |
| rows, perf = eval_pnsr(model, args, dev) | |
| else: | |
| assert args.intervention in ("none", "jobs_zero", "jself_zero"), \ | |
| "TX supports cache ablations only" | |
| rows, perf = eval_tx(model, args, dev) | |
| agg = aggregate(rows) | |
| agg["_perf"] = perf | |
| agg["_ci_memhard_beyond48"] = boot_ci(rows, lambda r: memhard(r) and r["delay"] > 48) | |
| agg["_ci_memhard"] = boot_ci(rows, memhard) | |
| agg["_n_rows"] = len(rows) | |
| tag = args.tag or f"{args.run}_{args.split}_{args.intervention}" + \ | |
| (f"_K{args.K}" if args.K else "") | |
| out = eval_root() / f"R_{tag}.json" | |
| atomic_write_json(out, agg) | |
| np.savez_compressed(eval_root() / f"rows_{tag}.npz", | |
| **{k: np.array([r[k] for r in rows]) for k in rows[0]}) | |
| print(json.dumps({k: v for k, v in agg.items() if not k.startswith("fam/")}, | |
| indent=1)[:1500]) | |
| print("->", out) | |
| if __name__ == "__main__": | |
| main() | |