#!/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()