File size: 12,970 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
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
#!/usr/bin/env python3
"""Attack the benchmark with shortcut nulls BEFORE training (preregistered).

NOTE: H is one-sided. H >= +0.05 is a leak; NEGATIVE H means the heuristic
performs BELOW its baseline, which is evidence of resistance, not of a
shortcut. Any gate must test H >= threshold, never |H| >= threshold.

Nulls (all restricted to model-visible inputs):
  enum families:
    uniform over legal set; majority (train prior, legal-masked);
    BoW logistic on query token counts (surface leakage detector).
  pointer families (per live-candidate at query):
    uniform over live candidates;
    metadata GBM (store/kind/key/ent-bucket/ages/ranks/counts - no text);
    lexical overlap (query tokens vs record key+val tokens), tie-break newest;
    newest-record; newest-in-store.
  op family: BoW -> op id (surface task, expected high; reported not gated).

Reported per family, with the corrected/uncorrected split for pointer families.
Gate (preregistration/PREREGISTRATION_PNS.md): H = (acc-base)/(1-base) < 0.05 for the
metadata GBM on every history family, and lexical-newest H < 0.10 on the
corrected pointer subfamily.  BoW leakage is reported and bounded by the
trained E-only control at eval time.
"""
import argparse
import json
import sys
from collections import Counter, defaultdict
from pathlib import Path

import numpy as np

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.data.view import Shard, shard_paths  # noqa: E402
from pns.world.schema import ENUM_VOCAB, Fam, Mode  # noqa: E402

POINTER = {int(f) for f in (Fam.EXACT_DELAYED, Fam.EXACT_2HOP, Fam.SELF_REF,
                            Fam.HANDLE_REF, Fam.GOAL_TOP)}
ENUMF = {int(f) for f in (Fam.SEM_LATEST, Fam.SEM_2HOP, Fam.TEMPORAL_ORDER,
                          Fam.DEADLINE, Fam.IMMEDIATE_CMP)}


def legal_ids(mask: np.ndarray) -> list[int]:
    out = []
    for i in range(len(ENUM_VOCAB)):
        if mask[i >> 3] & (1 << (i & 7)):
            out.append(i)
    return out


def iter_questions(paths, limit=None):
    n = 0
    for p in paths:
        sh = Shard(p)
        for i in range(sh.n_lifetimes):
            lv = sh.lifetime(i)
            recs = lv.records()
            for ev in lv:
                if ev.mode_gold in (int(Mode.ANSWER_POINTER), int(Mode.ANSWER_ENUM),
                                    int(Mode.EXTERNAL_OPERATION)):
                    yield ev, recs
                    n += 1
                    if limit and n >= limit:
                        return


def cand_features(ev, recs):
    """Metadata-only per-candidate features for pointer questions."""
    live = [(s, int(r)) for s, r in enumerate(ev.live_slots) if r >= 0]
    rows, is_gold = [], []
    b_all = np.array([int(_get(recs, "birth_ev", r)) for _, r in live])
    order = np.argsort(-b_all)  # newest first
    rank_of = {live[j][0]: int(np.where(order == j)[0][0]) + 1 for j in range(len(live))}
    keyent = Counter((_get(recs, "key", r), _get(recs, "ent", r)) for _, r in live)
    # rank within same (key, ent)
    for slot, r in live:
        k, e = _get(recs, "key", r), _get(recs, "ent", r)
        same = [(s2, r2) for s2, r2 in live
                if _get(recs, "key", r2) == k and _get(recs, "ent", r2) == e]
        same_sorted = sorted(same, key=lambda x: -_get(recs, "birth_ev", x[1]))
        same_rank = 1 + [s2 for s2, _ in same_sorted].index(slot)
        rows.append([
            _get(recs, "store", r), _get(recs, "kind", r), k, e,
            ev.idx - _get(recs, "birth_ev", r), rank_of[slot], same_rank,
            keyent[(k, e)], slot, len(live),
        ])
        is_gold.append(1 if slot == ev.ptr_gold_slot else 0)
    return np.asarray(rows, np.float32), np.asarray(is_gold, np.int8), live


def _get(recs, key, global_row):
    return int(recs[key][global_row - recs["_lo"]])


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--train-questions", type=int, default=60000)
    ap.add_argument("--val-questions", type=int, default=25000)
    ap.add_argument("--out", default=None)
    args = ap.parse_args()
    root = shards_root()
    train_paths = shard_paths("train", root)[:60]
    val_paths = shard_paths("val", root)

    # ---------- collect rows
    def collect(paths, limit):
        enum_rows = defaultdict(list)   # fam -> (qtok counts idx, legal, gold)
        ptr_rows = defaultdict(list)    # fam -> (X, y, corrected, n_live, lex_feats)
        op_rows = []
        for p in paths:
            sh = Shard(p)
            for i in range(sh.n_lifetimes):
                lv = sh.lifetime(i)
                recs = {k: v for k, v in lv.records().items()}
                recs["_lo"] = lv.rec_lo
                for ev in lv:
                    if ev.mode_gold == int(Mode.ANSWER_ENUM):
                        enum_rows[ev.family].append(
                            (ev.tokens.copy(), legal_ids(ev.enum_legal), ev.enum_gold))
                    elif ev.mode_gold == int(Mode.ANSWER_POINTER):
                        X, y, live = cand_features(ev, recs)
                        qset = set(ev.tokens.tolist())
                        lex = []
                        for slot, r in live:
                            row = r - lv.rec_lo
                            rset = set(recs["key_toks"][row].tolist()) | \
                                set(recs["val_toks"][row].tolist())
                            rset.discard(0)
                            lex.append((len(qset & rset) / max(1, len(rset)),
                                        int(recs["birth_ev"][row])))
                        g_row = [r for s, r in live if s == ev.ptr_gold_slot][0] - lv.rec_lo
                        g_store = int(recs["store"][g_row])
                        g_kind = int(recs["kind"][g_row])
                        g_key, g_ent = int(recs["key"][g_row]), int(recs["ent"][g_row])
                        n_store = n_kind = n_chain = 0
                        for _, r in live:
                            row = r - lv.rec_lo
                            if int(recs["store"][row]) == g_store:
                                n_store += 1
                                if int(recs["kind"][row]) == g_kind:
                                    n_kind += 1
                            if int(recs["key"][row]) == g_key and int(recs["ent"][row]) == g_ent:
                                n_chain += 1
                        ptr_rows[ev.family].append(
                            (X, y, ev.meta["corrected"], ev.meta["reverted"],
                             len(live), n_store, lex, n_kind, n_chain))
                    elif ev.mode_gold == int(Mode.EXTERNAL_OPERATION):
                        op_rows.append((ev.tokens.copy(), ev.op_gold))
                total = (sum(len(v) for v in enum_rows.values())
                         + sum(len(v) for v in ptr_rows.values()) + len(op_rows))
                if total >= limit:
                    return enum_rows, ptr_rows, op_rows
        return enum_rows, ptr_rows, op_rows

    print("collecting train rows ...", flush=True)
    tr_enum, tr_ptr, tr_op = collect(train_paths, args.train_questions)
    print("collecting val rows ...", flush=True)
    va_enum, va_ptr, va_op = collect(val_paths, args.val_questions)

    report = {"n_train": {}, "n_val": {}, "families": {}}

    # ---------- enum nulls
    from sklearn.linear_model import LogisticRegression
    from scipy.sparse import csr_matrix

    def bow(rows):
        data, indices, indptr = [], [], [0]
        for toks, _, _ in rows:
            c = Counter(toks.tolist())
            indices.extend(c.keys())
            data.extend(c.values())
            indptr.append(len(indices))
        return csr_matrix((data, indices, indptr), shape=(len(rows), 8192))

    for fam, rows in sorted(va_enum.items()):
        name = Fam(fam).name
        trows = tr_enum.get(fam, [])
        golds = np.array([g for _, _, g in rows])
        legal_sizes = np.array([len(l) for _, l, _ in rows])
        base = float(np.mean(1.0 / legal_sizes))
        out = {"n": len(rows), "base_uniform": round(base, 4)}
        prior = Counter(g for _, _, g in trows)
        maj_acc = float(np.mean([max(((prior.get(c, 0), c) for c in leg))[1] == g
                                 for _, leg, g in rows]))
        out["majority"] = round(maj_acc, 4)
        if trows:
            Xtr, Xva = bow(trows), bow(rows)
            ytr = np.array([g for _, _, g in trows])
            clf = LogisticRegression(max_iter=300, C=1.0, n_jobs=8)
            clf.fit(Xtr, ytr)
            proba = clf.predict_proba(Xva)
            classes = clf.classes_
            acc = 0
            for j, (_, leg, g) in enumerate(rows):
                mask = np.isin(classes, leg)
                if mask.sum() == 0:
                    continue
                pred = classes[mask][np.argmax(proba[j][mask])]
                acc += int(pred == g)
            bow_acc = acc / len(rows)
            out["bow_logistic"] = round(bow_acc, 4)
            out["H_bow"] = round((bow_acc - base) / (1 - base + 1e-9), 4)
            out["H_majority"] = round((maj_acc - base) / (1 - base + 1e-9), 4)
        report["families"][name] = out

    # ---------- pointer nulls
    from sklearn.ensemble import HistGradientBoostingClassifier

    for fam, rows in sorted(va_ptr.items()):
        name = Fam(fam).name
        trows = tr_ptr.get(fam, [])
        base = float(np.mean([1.0 / r[4] for r in rows]))
        base_store = float(np.mean([1.0 / r[5] for r in rows]))
        rev_rows = [r for r in rows if r[3] == 1]
        rev_base_store = float(np.mean([1.0 / r[5] for r in rev_rows])) if rev_rows else None
        out = {"n": len(rows), "n_reverted": len(rev_rows),
               "base_uniform": round(base, 4), "base_store": round(base_store, 4),
               "base_store_reverted": round(rev_base_store, 4) if rev_rows else None,
               "base_kind": round(float(np.mean([1.0 / r[7] for r in rows])), 4),
               "base_chain": round(float(np.mean([1.0 / r[8] for r in rows])), 4)}
        if rev_rows:
            out["base_chain_reverted"] = round(
                float(np.mean([1.0 / r[8] for r in rev_rows])), 4)

        def H(acc, b):
            return round((acc - b) / (1 - b + 1e-9), 4)

        def heur(pick):
            hits = np.array([int(r[1][pick(r[0], r[6])] == 1) for r in rows])
            rev = np.array([int(r[3] == 1) for r in rows], bool)
            a = float(hits.mean())
            ar = float(hits[rev].mean()) if rev.any() else None
            return a, ar

        a, ar = heur(lambda X, lex: int(np.argmin(X[:, 4])))
        out["newest"], out["newest_reverted"] = round(a, 4), \
            (round(ar, 4) if ar is not None else None)
        a, ar = heur(lambda X, lex: max(range(len(lex)),
                                        key=lambda j: (lex[j][0], lex[j][1])))
        out["lexical_newest"] = round(a, 4)
        out["lexical_newest_reverted"] = round(ar, 4) if ar is not None else None
        if ar is not None and out.get("base_chain_reverted"):
            out["H_lexical_newest_reverted_vs_chain"] = H(ar, out["base_chain_reverted"])
        if trows:
            Xtr = np.concatenate([r[0] for r in trows])
            ytr = np.concatenate([r[1] for r in trows])
            gbm = HistGradientBoostingClassifier(max_iter=200, max_depth=6)
            gbm.fit(Xtr, ytr)
            hits = np.array([int(r[1][int(np.argmax(gbm.predict_proba(r[0])[:, 1]))] == 1)
                             for r in rows])
            rev = np.array([int(r[3] == 1) for r in rows], bool)
            g_acc = float(hits.mean())
            out["metadata_gbm"] = round(g_acc, 4)
            out["H_metadata_gbm_vs_store"] = H(g_acc, base_store)
            out["H_metadata_gbm_vs_kind"] = H(g_acc, out["base_kind"])
            if rev.any():
                gr = float(hits[rev].mean())
                out["metadata_gbm_reverted"] = round(gr, 4)
                out["H_metadata_gbm_reverted_vs_chain"] = H(gr, out["base_chain_reverted"])
        report["families"][name] = out

    # ---------- op surface null
    if va_op and tr_op:
        Xtr, Xva = bow([(t, None, g) for t, g in tr_op]), bow([(t, None, g) for t, g in va_op])
        ytr = np.array([g for _, g in tr_op])
        clf = LogisticRegression(max_iter=300, n_jobs=8).fit(Xtr, ytr)
        acc = float(np.mean(clf.predict(Xva) == np.array([g for _, g in va_op])))
        report["families"]["OP_EMIT_opid"] = {
            "n": len(va_op), "bow_logistic": round(acc, 4),
            "note": "op named in request text; surface-solvable by design (not gated)"}

    out_path = Path(args.out) if args.out else eval_root() / "NULLS.json"
    atomic_write_json(out_path, report)
    print(json.dumps(report, indent=1))
    print("->", out_path)


if __name__ == "__main__":
    main()