File size: 12,485 Bytes
f930dac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
#!/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()