summaryrefslogtreecommitdiff
path: root/experiments/analyze_oral_a_full.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:25:06 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:25:06 -0500
commit94fc8471c515fb0baacc0aa79d4a47f539858e57 (patch)
tree39236544cef50f50559669a4bbf5bebd231c1b2e /experiments/analyze_oral_a_full.py
parent931273085034d48db5d93b71669ba9bc30b8bc4d (diff)
experiments: freeze oral A confirmation gates
Diffstat (limited to 'experiments/analyze_oral_a_full.py')
-rwxr-xr-xexperiments/analyze_oral_a_full.py83
1 files changed, 83 insertions, 0 deletions
diff --git a/experiments/analyze_oral_a_full.py b/experiments/analyze_oral_a_full.py
new file mode 100755
index 0000000..f22d55c
--- /dev/null
+++ b/experiments/analyze_oral_a_full.py
@@ -0,0 +1,83 @@
+#!/usr/bin/env python3
+"""Apply the frozen A3 full seed-0 ResNet-20 advancement gate."""
+import argparse
+import json
+import math
+import os
+
+
+SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b"
+
+
+def read_result(path, mode):
+ with open(path) as handle:
+ record = json.load(handle)
+ args = record["args"]
+ expected = {
+ "mode": mode, "depth": 20, "width": 16, "seed": 0,
+ "epochs": 200, "val_examples": 5000,
+ "eval_split": "validation", "normalization": "batchnorm",
+ }
+ for key, value in expected.items():
+ if args.get(key) != value:
+ raise ValueError(f"{path}: {key} drift")
+ if record["provenance"]["git_tracked_dirty"]:
+ raise ValueError(f"tracked-dirty result: {path}")
+ if record["split"]["validation_index_sha256"] != SPLIT_HASH:
+ raise ValueError(f"split drift: {path}")
+ if record["evaluation_protocol"]["test_evaluations"]:
+ raise ValueError(f"A3 touched test: {path}")
+ alignment = None
+ if record["diagnostics"] is not None:
+ alignment = float(record["diagnostics"]["early_third_mean"])
+ accuracy = float(record["final"]["accuracy"])
+ loss = float(record["final"]["loss"])
+ return {
+ "mode": mode, "path": path, "accuracy": accuracy, "loss": loss,
+ "finite": record["final"]["finite"] and math.isfinite(accuracy + loss),
+ "early_third_alignment": alignment,
+ "total_macs": record["work"]["total_macs_estimate"],
+ "peak_memory": record["hardware"]["peak_memory_allocated_bytes"],
+ "source_commit": record["provenance"]["git_commit"],
+ }
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--bp_selection", default="results/oral_a_bp_selection.json")
+ parser.add_argument("--dfa", default="results/oral_a_dev/dfa_full_r20_s0.json")
+ parser.add_argument("--sdil", default="results/oral_a_dev/sdil_full_r20_s0.json")
+ parser.add_argument("--out", default="results/oral_a_full_gate.json")
+ args = parser.parse_args()
+ with open(args.bp_selection) as handle:
+ bp_selection = json.load(handle)
+ if not bp_selection["status"].startswith("passed_"):
+ raise ValueError("A1 BP reference is not viable")
+ bp = read_result(bp_selection["selected"]["path"], "bp")
+ dfa = read_result(args.dfa, "dfa")
+ sdil = read_result(args.sdil, "sdil")
+ checks = {
+ "bp_at_least_90": bp["accuracy"] >= 0.90,
+ "sdil_within_5pt_bp": sdil["accuracy"] >= bp["accuracy"] - 0.05,
+ "sdil_beats_dfa_2pt": sdil["accuracy"] >= dfa["accuracy"] + 0.02,
+ "sdil_alignment_at_least_0p05": (
+ sdil["early_third_alignment"] is not None
+ and sdil["early_third_alignment"] >= 0.05),
+ "all_finite": all(row["finite"] for row in (bp, dfa, sdil)),
+ "sdil_macs_no_more_than_bp": sdil["total_macs"] <= bp["total_macs"],
+ }
+ output = {
+ "protocol": "oral_a_A3_v1",
+ "status": "passed" if all(checks.values()) else "failed",
+ "checks": checks, "rows": [bp, dfa, sdil],
+ "confirmation_test_seeds_touched": False,
+ }
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ with open(args.out, "w") as handle:
+ json.dump(output, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps(output, indent=2))
+
+
+if __name__ == "__main__":
+ main()