#!/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 @torch.no_grad() 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)), )) @torch.no_grad() 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()