Download eval/decode_probe.py from nur-dev/pns-bind-25m: direct link, hf CLI and curl.
- Browser
- Download file 12.5 kB
-
https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/decode_probe.py
- Command line
-
hf download hf://nur-dev/pns-bind-25m/eval/decode_probe.py
-
curl -L -o decode_probe.py https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/decode_probe.py
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)] | |
| 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() | |