pns-bind-25m / eval /attack_nulls.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw
History Blame Contribute Delete
13 kB
#!/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()