pns-bind-25m / eval /make_tables.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
7.98 kB
#!/usr/bin/env python3
"""Emit the manuscript's result tables as CSV from `reproduced_headline.json`.
python3 eval/make_tables.py [--headline results/reproduced_headline.json]
Writes, under results/:
experiment1_causal.csv causal handles per state model
experiment1_accuracy.csv per-family accuracy per model
experiment2_stageA.csv Stage-A gate
experiment2_confirmation.csv sealed-split confirmation and its verdict
decode_probe.csv probe power control and durable-binding result
long_delay.csv accuracy vs horizon and vs evidence delay
Every number in the paper's main result tables is one row of one of these.
"""
import argparse
import csv
import json
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT / "src"))
from pns.common import eval_root # noqa: E402
FAMS = ["SEM_LATEST", "SEM_2HOP", "IMMEDIATE_CMP", "EXACT_DELAYED", "SELF_REF",
"HANDLE_REF", "DEADLINE", "GOAL_TOP", "all"]
MECHANISM = {
"PNSR_K4_fin": "additive write + norm clamp tau=16",
"PNSR_K4_s2": "additive write + norm clamp tau=16",
"PNSR_K4_s3": "additive write + norm clamp tau=16",
"G_TAU512": "additive write + norm clamp tau=512",
"G_CONVEX_FREE": "convex gated update, no clamp",
"G_CONVEX_RETAIN": "convex gated update, retention-biased",
"RMT_s1": "memory tokens + full self-attention",
"PNSR_K1_fin": "additive write + norm clamp tau=16 (K=1)",
"TX768_fin": "none (768-token transcript window)",
"TXE_fin": "none (current event only)",
}
SEED = {"PNSR_K4_fin": 1, "PNSR_K4_s2": 2, "PNSR_K4_s3": 3}
def w(path, header, rows):
with open(path, "w", newline="") as f:
c = csv.writer(f)
c.writerow(header)
c.writerows(rows)
print(f" {path.name:32s} {len(rows)} rows")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--headline", default=str(ROOT / "results" / "reproduced_headline.json"))
ap.add_argument("--outdir", default=str(ROOT / "results"))
args = ap.parse_args()
H = json.loads(Path(args.headline).read_text())
out = Path(args.outdir)
out.mkdir(parents=True, exist_ok=True)
e1 = H["experiment1_generic_recurrent_state"]["models"]
# -------------------------------------------------- experiment1_causal.csv
rows = []
for run, m in e1.items():
c = m.get("learned_state_causal", {})
les = m.get("exact_store_lesions", {})
st = m.get("state_statistics", {})
rows.append([run, SEED.get(run, 1), MECHANISM.get(run, ""),
c.get("memhard_warm_val_1k"),
c.get("delta_cross_lifetime_swap"), c.get("delta_reset64"),
c.get("delta_reset32"), c.get("delta_state_zero"),
les.get("J_self_lesion_SELF_REF"),
les.get("J_obs_lesion_pointer_families"),
st.get("mean_cross_lifetime_cosine"), st.get("max_slot_norm")])
w(out / "experiment1_causal.csv",
["run", "seed", "state_mechanism", "memhard_acc_warm_L1024",
"delta_cross_lifetime_swap", "delta_reset64", "delta_reset32",
"delta_state_zero", "delta_J_self_lesion_SELF_REF",
"delta_J_obs_lesion_pointer_families", "mean_cross_lifetime_cosine",
"max_slot_norm"], rows)
# ------------------------------------------------ experiment1_accuracy.csv
rows = []
for run, m in e1.items():
a = m["semantic_accuracy"]
for fam in FAMS:
if a.get(fam) is not None:
rows.append([run, MECHANISM.get(run, ""), "val", 256, fam, a[fam]])
w(out / "experiment1_accuracy.csv",
["run", "state_mechanism", "split", "lifetime_length", "family", "accuracy"], rows)
# ----------------------------------------------------------- long_delay.csv
rows = []
for run, m in e1.items():
for horizon, d in (m.get("horizons") or {}).items():
for k, v in d.items():
if v is not None:
rows.append([run, "experiment1", horizon, k, v])
e2 = H.get("experiment2_protected_bindings") or {}
for arm in ("bind", "capacity_matched_unbound"):
for seed, d in (e2.get(arm) or {}).items():
for k in ("accuracy_delay_gt16", "accuracy_delay_gt128", "overall"):
if d.get(k) is not None:
rows.append([f"E3B_{arm}_{seed}", "experiment2", e2["split"], k, d[k]])
w(out / "long_delay.csv",
["run", "experiment", "horizon_or_split", "metric", "value"], rows)
# -------------------------------------------------------- decode_probe.csv
rows = []
for label, d in (H.get("experiment1_decode_probe") or {}).items():
for k, v in d.items():
rows.append([label, d.get("split_train"), d.get("split_test"), k, v])
w(out / "decode_probe.csv",
["label", "decoder_train_split", "decoder_test_split", "metric", "value"], rows)
# ------------------------------------------------ experiment2_stageA.csv
p = eval_root() / "STAGEA_GATE.json"
rows = []
if p.exists():
j = json.loads(p.read_text())
if "bind" in j and isinstance(j["bind"], dict) and "1" in j["bind"]:
# the schema the study's own Stage-A reporter wrote; kept readable
# so results/published/stage_a_gate_output.json still parses
for mode in ("bind", "unbound"):
for s, r in sorted(j[mode].items()):
rows.append([f"E3A_{mode}_s{s}", mode, int(s), r["acc"],
r["gt16"], r["gt128"], None, r["swap"],
r["zero"], r["meets_all"]])
else:
for run, r in j.items():
if run == "verdict":
continue
rows.append([run, r["mode"], r["seed"], round(r["none"]["acc"], 4),
r["none"].get("acc_gt16"), r["none"].get("acc_gt128"),
r["none"].get("n_gt128"), r["delta_payload_swap"],
r["delta_payload_zero"], r["meets_all"]])
rows.sort(key=lambda x: (x[1], x[2]))
w(out / "experiment2_stageA.csv",
["run", "arm", "seed", "accuracy", "accuracy_delay_gt16",
"accuracy_delay_gt128", "n_delay_gt128", "delta_payload_swap",
"delta_payload_zero", "meets_gate"], rows)
# ------------------------------------------ experiment2_confirmation.csv
rows = []
if e2:
t = e2["TXE_memory_free_control"]
rows.append(["TXE_ON_E3", "memory_free_control", "-", t["latest"],
t["2hop"], t["pooled"], "", "", "", "", "", ""])
for arm, label in (("bind", "binding_local_writes"),
("capacity_matched_unbound", "capacity_matched_unstructured")):
for seed, d in (e2.get(arm) or {}).items():
rows.append([f"E3B_{'bind' if arm == 'bind' else 'unbound'}_{seed}",
label, seed, d["SEM_LATEST"], d["SEM_2HOP"], d["overall"],
d["accuracy_delay_gt128"], d["delta_payload_swap"],
d["delta_reset64"], d["delta_payload_zero"],
d["accuracy_under_payload_zero"],
e2["verdict"].get(seed, {}).get("all", "")])
w(out / "experiment2_confirmation.csv",
["run", "arm", "seed", "SEM_LATEST", "SEM_2HOP", "overall",
"accuracy_delay_gt128", "delta_payload_swap", "delta_reset64",
"delta_payload_zero", "accuracy_under_payload_zero",
"conjunction_met_this_seed"], rows)
if e2:
v = e2["verdict"]
print(f"\n CONJUNCTION_MET = {v['CONJUNCTION_MET']} "
f"(pooled gain {v['pooled_gain']}, "
f"SEM_2HOP endpoint met = {e2['preregistered_endpoint_SEM_2HOP']['met']})")
if __name__ == "__main__":
main()