pns-bind-25m / eval /evaluate.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
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
@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()