pns-bind-25m / eval /stage_a_eval.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
3.04 kB
#!/usr/bin/env python3
"""Experiment 2 Stage-A gate, evaluated from published checkpoints.
Gate frozen in preregistration/PREREGISTRATION_E3_BIND.md:
acc(delay > 16) >= 0.30, acc(delay > 128) >= 0.25,
delta payload-swap >= 0.10, delta payload-zero >= 0.10
all three seeds positive, at least two meeting every number.
python3 eval/stage_a_eval.py
"""
import argparse
import json
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT / "src"))
sys.path.insert(0, str(Path(__file__).resolve().parent))
from pns.checkpoint import load_model # noqa: E402
from pns.common import atomic_write_json, eval_root # noqa: E402
from stage_a import evaluate, summarise # noqa: E402
GATE = dict(acc_gt16=0.30, acc_gt128=0.25, delta_payload_swap=0.10,
delta_payload_zero=0.10)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--split", default="e3_dev")
ap.add_argument("--eval-lifetimes", type=int, default=400)
args = ap.parse_args()
dev = "cuda"
rep, table = {}, []
for mode in ("bind", "unbound"):
for s in (1, 2, 3):
run = f"E3A_{mode}_s{s}"
model, _, _ = load_model(run, dev)
r = {"run": run, "mode": mode, "seed": s}
for iv in ("none", "payload_swap", "payload_zero"):
r[iv] = summarise(evaluate(model, dev, args.split,
args.eval_lifetimes, iv))
r["delta_payload_swap"] = round(r["none"]["acc"] - r["payload_swap"]["acc"], 4)
r["delta_payload_zero"] = round(r["none"]["acc"] - r["payload_zero"]["acc"], 4)
vals = {"acc_gt16": r["none"].get("acc_gt16", 0.0),
"acc_gt128": r["none"].get("acc_gt128", 0.0),
"delta_payload_swap": r["delta_payload_swap"],
"delta_payload_zero": r["delta_payload_zero"]}
r["meets_all"] = all(vals[k] >= GATE[k] for k in GATE)
rep[run] = r
table.append((mode, s, round(r["none"]["acc"], 4), vals["acc_gt16"],
vals["acc_gt128"], vals["delta_payload_swap"],
vals["delta_payload_zero"], r["meets_all"]))
print(run, json.dumps(vals), "meets_all", r["meets_all"], flush=True)
def verdict(mode):
rs = [rep[f"E3A_{mode}_s{s}"] for s in (1, 2, 3)]
pos = all(r["delta_payload_swap"] > 0 and r["delta_payload_zero"] > 0 for r in rs)
return bool(pos and sum(r["meets_all"] for r in rs) >= 2)
rep["verdict"] = {"bind_passes": verdict("bind"), "unbound_passes": verdict("unbound")}
print("\n| arm | seed | acc | acc d>16 | acc d>128 | d swap | d zero | meets all |")
print("| --- | --- | --- | --- | --- | --- | --- | --- |")
for row in table:
print("| " + " | ".join(str(x) for x in row) + " |")
print("\nverdict:", rep["verdict"])
atomic_write_json(eval_root() / "STAGEA_GATE.json", rep)
if __name__ == "__main__":
main()