summaryrefslogtreecommitdiff
path: root/experiments/analyze_oral_a_apical.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:22:25 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:22:25 -0500
commit931273085034d48db5d93b71669ba9bc30b8bc4d (patch)
treef0b5da7681f9feefb9fa7b3fdd77e8ea4c78d1af /experiments/analyze_oral_a_apical.py
parentc6af287a6177cda3863e6f25f2bdab8abd611ad8 (diff)
experiments: freeze oral A development gates
Diffstat (limited to 'experiments/analyze_oral_a_apical.py')
-rwxr-xr-xexperiments/analyze_oral_a_apical.py100
1 files changed, 100 insertions, 0 deletions
diff --git a/experiments/analyze_oral_a_apical.py b/experiments/analyze_oral_a_apical.py
new file mode 100755
index 0000000..1d10edb
--- /dev/null
+++ b/experiments/analyze_oral_a_apical.py
@@ -0,0 +1,100 @@
+#!/usr/bin/env python3
+"""Apply the frozen A2a selector to convolutional apical-screen records."""
+import argparse
+import glob
+import json
+import math
+import os
+
+
+MODES = ("spatial_template", "channel_gated")
+EXPECTED_SCALES = (0.0, 1.0)
+EXPECTED_RATES = (0.001, 0.01, 0.1)
+
+
+def load(path):
+ with open(path) as handle:
+ record = json.load(handle)
+ args = record["args"]
+ if record["provenance"]["git_tracked_dirty"]:
+ raise ValueError(f"tracked-dirty result: {path}")
+ expected = {
+ "mode": "sdil", "depth": 20, "width": 16, "seed": 0,
+ "epochs": 0, "train_limit": 10000, "val_examples": 5000,
+ "a_warmup_steps": 100, "pert_directions": 1, "pert_every": 4,
+ "normalization": "batchnorm",
+ }
+ for key, value in expected.items():
+ if args.get(key) != value:
+ raise ValueError(f"{path}: {key}={args.get(key)!r}, expected {value!r}")
+ if record["split"]["validation_index_sha256"] != (
+ "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b"):
+ raise ValueError(f"split drift: {path}")
+ diagnostics = record.get("diagnostics")
+ warmup = record.get("apical_warmup", {}).get("last")
+ if diagnostics is None or warmup is None:
+ raise ValueError(f"missing diagnostics/warmup: {path}")
+ values = diagnostics["teaching_negative_gradient_cosine"]
+ early_count = max(1, len(values) // 3)
+ metrics = {
+ "early_third_alignment": sum(values[:early_count]) / early_count,
+ "all_layer_alignment": sum(values) / len(values),
+ "calibration_mse": warmup["calibration_mse"],
+ "prediction_target_cosine": warmup["prediction_target_cosine"],
+ }
+ finite = (record["final"]["finite"]
+ and all(math.isfinite(value) for value in metrics.values()))
+ return {
+ "path": path,
+ "source_commit": record["provenance"]["git_commit"],
+ "vectorizer_mode": args["vectorizer_mode"],
+ "a_scale": float(args["a_scale"]),
+ "eta_A": float(args["eta_A"]),
+ "metrics": metrics,
+ "eligible": finite and metrics["early_third_alignment"] > 0.0,
+ "vectorizer_parameters": record["architecture"]["vectorizer_parameters"],
+ }
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--input", default="results/oral_a_apical_screen")
+ parser.add_argument("--out", default="results/oral_a_apical_selection.json")
+ args = parser.parse_args()
+ paths = sorted(glob.glob(os.path.join(args.input, "*.json")))
+ rows = [load(path) for path in paths]
+ observed = {(row["vectorizer_mode"], row["a_scale"], row["eta_A"])
+ for row in rows}
+ expected = {(mode, scale, rate) for mode in MODES
+ for scale in EXPECTED_SCALES for rate in EXPECTED_RATES}
+ if observed != expected or len(rows) != len(expected):
+ raise ValueError(f"incomplete A2a grid: missing={expected-observed}, extra={observed-expected}")
+ if len({row["source_commit"] for row in rows}) != 1:
+ raise ValueError("A2a source commits differ")
+ selected = {}
+ for mode in MODES:
+ eligible = [row for row in rows
+ if row["vectorizer_mode"] == mode and row["eligible"]]
+ if eligible:
+ eligible.sort(key=lambda row: (
+ -row["metrics"]["early_third_alignment"],
+ -row["metrics"]["all_layer_alignment"],
+ row["eta_A"], row["a_scale"]))
+ selected[mode] = eligible[0]
+ output = {
+ "protocol": "oral_a_A2a_v1",
+ "status": ("selected" if len(selected) == len(MODES)
+ else "failed_no_eligible_family"),
+ "rows": rows,
+ "selected": selected,
+ "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({"status": output["status"], "selected": selected}, indent=2))
+
+
+if __name__ == "__main__":
+ main()