#!/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()