pns-bind-25m / eval /decode_probe.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
12.5 kB
#!/usr/bin/env python3
"""Is lifetime-specific information PRESENT in Sigma but unreadable, or ABSENT?
The interventions show Sigma is causally inert; the probes show no transition
rule restores long-delay retention. Two very different mechanisms remain:
(A) information IS written into Sigma but is overwritten/entangled by later
events so the model's readout cannot recover it -> a decoder/organisation
problem, and a strong external decoder should still find it;
(B) information is destroyed (or never durably written) -> no decoder can
recover it, and the failure is in the state dynamics themselves.
This script decides between them. It freezes a trained model, harvests Sigma
at SEM_LATEST question events, and trains an EXTERNAL decoder (multinomial
logistic regression) on the frozen state to predict the gold answer, scored
under the same legal-set mask the model uses. Reported by delay bucket.
Controls, all required for the result to mean anything:
shuffled - Sigma rows permuted across examples: measures what the decoder
can get from label priors alone (its capacity floor).
event - decode from the current-event vector instead of Sigma: for
delayed questions this should be near chance, and it bounds how
much of any Sigma result is really "the question is in the text".
"""
import argparse
import sys
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.train.loader import iter_eval_batches # noqa: E402
from pns.world.schema import ENUM_VOCAB, Fam # noqa: E402
from evaluate import load_model, to_gpu # noqa: E402
BUCKETS = [(1, 8), (9, 16), (17, 32), (33, 64), (65, 192)]
@torch.no_grad()
def harvest(model, dev, n_lifetimes, split):
"""Collect (Sigma_{t-}, query vector, gold, legal-mask, delay, lifetime id)
at SEM_LATEST questions. Sigma is snapshotted BEFORE the question event is
processed, so the question never contributes to what is decoded."""
S, V, Y, Lg, Dl, LT, PE = [], [], [], [], [], [], []
d = model.cfg.d
for b in iter_eval_batches(split, shards_root(), 48, n_lifetimes):
g = to_gpu(b, dev)
B, L = g["etype"].shape
state = model.initial_state(B, dev)
with torch.autocast("cuda", dtype=torch.bfloat16):
h = 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):
live = g["live"][:, t]; mask = live >= 0; rows = live.clamp(min=0)
bank = torch.gather(h, 1, rows.unsqueeze(-1).expand(-1, -1, d))
bank = model.recenc.finalize(
bank, (t - torch.gather(g["rec_birth"], 1, rows)).clamp(min=0))
prev = state
state, _ = model.step(state, g["tok"][:, t], g["etype"][:, t],
g["dt"][:, t], bank * mask.unsqueeze(-1), mask)
sel = g["family"][:, t] == int(Fam.SEM_LATEST)
if sel.any():
# state BEFORE the question's own write: the question must
# be answered from history, not from itself
idx = torch.nonzero(sel).flatten()
S.append(prev[idx].float().cpu().numpy())
_, v, _ = model.encode_event(g["tok"][idx, t], g["etype"][idx, t],
g["dt"][idx, t])
V.append(v.float().cpu().numpy())
Y.append(g["enum_gold"][idx, t].cpu().numpy())
Lg.append(g["enum_legal"][idx, t].cpu().numpy())
Dl.append(g["delay"][idx, t].cpu().numpy())
LT.append(g["seed"][idx].cpu().numpy())
PE.append(g["etype"][idx, max(t - 1, 0)].cpu().numpy())
return (np.concatenate(S), np.concatenate(V), np.concatenate(Y),
np.concatenate(Lg), np.concatenate(Dl), np.concatenate(LT),
np.concatenate(PE))
def masked_acc(proba, classes, Y, Lg):
ok = 0
for j in range(len(Y)):
legal = [i for i in range(len(ENUM_VOCAB))
if Lg[j][i >> 3] & (1 << (i & 7))]
m = np.isin(classes, legal)
if not m.any():
continue
ok += int(classes[m][np.argmax(proba[j][m])] == Y[j])
return ok / max(len(Y), 1)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--run", default="PNSR_K4_fin")
ap.add_argument("--train-lifetimes", type=int, default=256)
ap.add_argument("--test-lifetimes", type=int, default=128)
ap.add_argument("--refresh", action="store_true")
ap.add_argument("--pca", type=int, default=256)
ap.add_argument("--train-split", default="train")
ap.add_argument("--val-split", default="val")
ap.add_argument("--tag", default=None,
help="output suffix; defaults to the run name")
args = ap.parse_args()
dev = "cuda"
model, _ = load_model(args.run, dev)
tag = args.tag or args.run
cache = eval_root() / f"DECODE_CACHE_{tag}.npz"
if cache.exists() and not args.refresh:
z = np.load(cache)
Str, Vtr, Ytr, Ltr, Dtr, LTtr, PEtr = (z[k] for k in
("Str", "Vtr", "Ytr", "Ltr", "Dtr", "LTtr", "PEtr"))
Ste, Vte, Yte, Lte, Dte, LTte, PEte = (z[k] for k in
("Ste", "Vte", "Yte", "Lte", "Dte", "LTte", "PEte"))
print("loaded cached harvest", cache, flush=True)
else:
print(f"harvesting train={args.train_split} test={args.val_split} ...",
flush=True)
Str, Vtr, Ytr, Ltr, Dtr, LTtr, PEtr = harvest(
model, dev, args.train_lifetimes, args.train_split)
Ste, Vte, Yte, Lte, Dte, LTte, PEte = harvest(
model, dev, args.test_lifetimes, args.val_split)
np.savez(cache, Str=Str, Vtr=Vtr, Ytr=Ytr, Ltr=Ltr, Dtr=Dtr, LTtr=LTtr,
PEtr=PEtr, Ste=Ste, Vte=Vte, Yte=Yte, Lte=Lte, Dte=Dte,
LTte=LTte, PEte=PEte)
# train/test come from disjoint splits (train shards vs val shards), so the
# lifetime-grouping requirement is satisfied by construction; assert it.
assert not (set(LTtr.tolist()) & set(LTte.tolist())), "lifetime leak"
print(f"train {Str.shape} test {Ste.shape} "
f"(disjoint lifetimes: {len(set(LTtr.tolist()))}/{len(set(LTte.tolist()))})",
flush=True)
from sklearn.linear_model import LogisticRegression
from sklearn.neural_network import MLPClassifier
from sklearn.preprocessing import StandardScaler
rep = {"run": args.run, "n_train": int(len(Ytr)), "n_test": int(len(Yte)),
"train_split": args.train_split, "val_split": args.val_split,
"train_lifetimes": int(len(set(LTtr.tolist()))),
"test_lifetimes": int(len(set(LTte.tolist())))}
def score(clf, sc, Xte, tag):
p = clf.predict_proba(sc.transform(Xte))
out = {"overall": round(masked_acc(p, clf.classes_, Yte, Lte), 4)}
for lo, hi in BUCKETS:
m = (Dte >= lo) & (Dte <= hi)
if m.sum() > 30:
out[f"d{lo}-{hi}"] = round(
masked_acc(p[m], clf.classes_, Yte[m], Lte[m]), 4)
rep[tag] = out
print(f"{tag:34s}", out, flush=True)
return out
def fit_eval(Xtr, Xte, tag, kind):
sc = StandardScaler().fit(Xtr)
clf = (LogisticRegression(max_iter=2000, C=1.0, n_jobs=16) if kind == "linear"
else MLPClassifier(hidden_layer_sizes=(512,), max_iter=120,
early_stopping=True, random_state=0))
clf.fit(sc.transform(Xtr), Ytr)
score(clf, sc, Xte, f"{tag}::{kind}")
return clf, sc
from sklearn.decomposition import PCA
raw_tr, raw_te = Str.reshape(len(Str), -1), Ste.reshape(len(Ste), -1)
ncomp = min(args.pca, len(raw_tr) // 8, raw_tr.shape[1])
pca = PCA(n_components=ncomp, random_state=0).fit(raw_tr)
flat_tr, flat_te = pca.transform(raw_tr), pca.transform(raw_te)
rep["pca_components"] = int(ncomp)
rep["pca_explained_var"] = round(float(pca.explained_variance_ratio_.sum()), 4)
rep["examples_per_feature"] = round(len(raw_tr) / (ncomp + Vtr.shape[1]), 2)
print(f"PCA {ncomp} comps, {rep['pca_explained_var']:.3f} var, "
f"{rep['examples_per_feature']} examples/feature", flush=True)
# ---- PROBE-POWER POSITIVE CONTROL -------------------------------------
# Decode the PREVIOUS event's type from Sigma_{t-}. That event just wrote
# into Sigma under typed eligibility, so if Sigma encodes anything at all
# this is recoverable. If this control is at chance, the probe is
# underpowered and NO null from it is interpretable.
from sklearn.preprocessing import StandardScaler as _SS
from sklearn.neural_network import MLPClassifier as _MLP
sc0 = _SS().fit(flat_tr)
pc = _MLP(hidden_layer_sizes=(512,), max_iter=120, early_stopping=True,
random_state=0).fit(sc0.transform(flat_tr), PEtr)
acc_pc = float((pc.predict(sc0.transform(flat_te)) == PEte).mean())
vals, cnts = np.unique(PEtr, return_counts=True)
maj = float((PEte == vals[np.argmax(cnts)]).mean())
rep["power_control_prev_etype"] = {
"acc": round(acc_pc, 4), "majority_baseline": round(maj, 4),
"passes": bool(acc_pc > maj + 0.10)}
print("POWER CONTROL (decode previous event type from Sigma):",
rep["power_control_prev_etype"], flush=True)
# the decoder must be told WHAT to read: the question-event vector supplies
# entity/attribute/referent identity, never the historical answer
SQ_tr = np.concatenate([flat_tr, Vtr], 1)
SQ_te = np.concatenate([flat_te, Vte], 1)
# control: pair each label with ANOTHER LIFETIME's state, query kept correct
rng = np.random.default_rng(0)
def cross_lifetime_perm(LT):
idx = np.arange(len(LT))
for _ in range(200):
p_ = rng.permutation(idx)
if (LT[p_] != LT).mean() > 0.98:
return p_
return p_
def matched_perm(LT, D):
"""Cross-lifetime, matched on delay bucket, so the control cannot be
beaten via state-age <-> answer-distribution correlations."""
out = np.arange(len(LT))
bid = np.digitize(D, [9, 17, 33, 65])
for b in np.unique(bid):
m = np.where(bid == b)[0]
if len(m) < 2:
continue
best = None
for _ in range(200):
p_ = rng.permutation(m)
frac = float((LT[p_] != LT[m]).mean())
if best is None or frac > best[1]:
best = (p_, frac)
if frac > 0.98:
break
out[m] = best[0]
return out
ptr, pte = matched_perm(LTtr, Dtr), matched_perm(LTte, Dte)
rep["shuffle_cross_lifetime_frac"] = round(float((LTte[pte] != LTte).mean()), 4)
SH_tr = np.concatenate([flat_tr[ptr], Vtr], 1)
SH_te = np.concatenate([flat_te[pte], Vte], 1)
for kind in ("linear", "mlp"):
clf, sc = fit_eval(SQ_tr, SQ_te, "sigma+query", kind)
fit_eval(Vtr, Vte, "query_only", kind)
fit_eval(SH_tr, SH_te, "shuffled+query", kind)
# ANCILLARY: the SAME frozen intact-trained decoder, evaluated on
# cross-lifetime-swapped test states. Asks whether that decoder itself
# depended on lifetime-specific state, rather than what a separately
# retrained decoder could achieve without it.
intact = rep[f"sigma+query::{kind}"]
swapped = score(clf, sc, SH_te, f"paired_testtime_swap::{kind}")
rep[f"delta_probe_use::{kind}"] = {
k: round(intact[k] - swapped[k], 4) for k in intact if k in swapped}
print(f" delta_probe_use ({kind}):", rep[f"delta_probe_use::{kind}"],
flush=True)
rep["chance"] = round(float(np.mean([1.0 / max(sum(
1 for i in range(len(ENUM_VOCAB)) if Lte[j][i >> 3] & (1 << (i & 7))), 1)
for j in range(len(Yte))])), 4)
atomic_write_json(eval_root() / f"DECODE_SIGMA_{tag}.json", rep)
print("chance", rep["chance"], "->", eval_root() / f"DECODE_SIGMA_{tag}.json")
if __name__ == "__main__":
main()