#!/usr/bin/env python3 """Apply the frozen Oral-A-v2 causal-capture gate.""" import argparse import glob import json import math import os MODES = ("unit_targets", "channel_subspace") RATES = (0.01, 0.1, 1.0) SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b" def load(path): with open(path) as handle: record = json.load(handle) args = record["args"] expected = { "mode": "sdil", "depth": 20, "width": 16, "seed": 0, "epochs": 0, "train_limit": 10000, "val_examples": 5000, "a_warmup_steps": 400, "pert_directions": 1, "pert_every": 4, "pert_sigma": 0.01, "perturb_seed": 1000, "normalization": "batchnorm", "vectorizer_mode": "channel_gated", "a_scale": 0.0, "alignment_probe": 64, } 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["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}") mode = args["apical_calibration_mode"] expected_space = ("channel_basis_moments" if mode == "channel_subspace" else "full_hidden_field") if record.get("calibration_metric_space") != expected_space: raise ValueError(f"calibration metric-space drift: {path}") diagnostics = record.get("diagnostics") warmup = record.get("apical_warmup", {}).get("mean") if diagnostics is None or warmup is None: raise ValueError(f"missing diagnostics/warmup aggregate: {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), "mean_calibration_mse": warmup["calibration_mse"], "mean_target_power": warmup["target_power"], "mean_prediction_target_cosine": warmup["prediction_target_cosine"], "mean_parameter_update_rms": warmup.get("parameter_update_rms", 0.0), } finite = (record["final"]["finite"] and all(math.isfinite(value) for value in metrics.values())) return { "path": path, "source_commit": record["provenance"]["git_commit"], "calibration_mode": mode, "eta_A": float(args["eta_A"]), "metric_space": expected_space, "metrics": metrics, "finite": finite, } def main(): parser = argparse.ArgumentParser() parser.add_argument("--input", default="results/oral_a_v2_calibration") parser.add_argument("--out", default="results/oral_a_v2_calibration_gate.json") args = parser.parse_args() rows = [load(path) for path in sorted(glob.glob( os.path.join(args.input, "*.json")))] observed = {(row["calibration_mode"], row["eta_A"]) for row in rows} expected = {(mode, rate) for mode in MODES for rate in RATES} if observed != expected or len(rows) != len(expected): raise ValueError( f"incomplete v2 grid: missing={expected-observed}, extra={observed-expected}") if len({row["source_commit"] for row in rows}) != 1: raise ValueError("v2 calibration source commits differ") selected = {} for mode in MODES: candidates = [row for row in rows if row["calibration_mode"] == mode and row["finite"]] if candidates: candidates.sort(key=lambda row: ( -row["metrics"]["early_third_alignment"], -row["metrics"]["all_layer_alignment"], row["eta_A"])) selected[mode] = candidates[0] checks = { "all_six_records_finite": all(row["finite"] for row in rows), "both_modes_selected": len(selected) == len(MODES), } if checks["both_modes_selected"]: structured = selected["channel_subspace"]["metrics"] unit = selected["unit_targets"]["metrics"] checks.update({ "structured_early_third_at_least_0.01": ( structured["early_third_alignment"] >= 0.01), "structured_all_layer_at_least_0.01": ( structured["all_layer_alignment"] >= 0.01), "structured_early_gain_over_unit_at_least_0.01": ( structured["early_third_alignment"] - unit["early_third_alignment"] >= 0.01), }) else: checks.update({ "structured_early_third_at_least_0.01": False, "structured_all_layer_at_least_0.01": False, "structured_early_gain_over_unit_at_least_0.01": False, }) passed = all(checks.values()) output = { "protocol": "oral_a_v2_causal_capture_v1", "status": "passed" if passed else "failed", "checks": checks, "rows": rows, "selected": selected, "confirmation_test_seeds_touched": False, "review_score_before": 5, "review_score_after": 5, "score_change_rule": "mechanics/calibration alone cannot raise score", } 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"], "checks": checks, "selected": selected, }, indent=2)) if __name__ == "__main__": main()