summaryrefslogtreecommitdiff
path: root/experiments/analyze_kp_stability_margin.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:50:44 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:50:44 -0500
commit0ae690e7b3144da2c5b3384cd7169e86529fd26f (patch)
tree7b5416b1dca30598413fb7b4cabd6c2011484b1a /experiments/analyze_kp_stability_margin.py
parent293785af4c9a556f388f04bcf9599aaee9e9dfd9 (diff)
protocol: freeze stability-margin training screen
Diffstat (limited to 'experiments/analyze_kp_stability_margin.py')
-rw-r--r--experiments/analyze_kp_stability_margin.py175
1 files changed, 175 insertions, 0 deletions
diff --git a/experiments/analyze_kp_stability_margin.py b/experiments/analyze_kp_stability_margin.py
new file mode 100644
index 0000000..d9599c9
--- /dev/null
+++ b/experiments/analyze_kp_stability_margin.py
@@ -0,0 +1,175 @@
+#!/usr/bin/env python3
+"""Audit and select the frozen S0 training-only stability margin."""
+import argparse
+import json
+import math
+import os
+import statistics
+
+
+MARGINS = (
+ ("0p001", 0.001),
+ ("0p003", 0.003),
+ ("0p01", 0.01),
+ ("0p03", 0.03),
+)
+SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b"
+
+
+def numeric_leaves(value):
+ if isinstance(value, bool) or value is None:
+ return
+ if isinstance(value, (int, float)):
+ yield float(value)
+ elif isinstance(value, dict):
+ for child in value.values():
+ yield from numeric_leaves(child)
+ elif isinstance(value, (list, tuple)):
+ for child in value:
+ yield from numeric_leaves(child)
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--input_dir", default="results/kp_innovation_stability_grid")
+ parser.add_argument(
+ "--out", default="results/kp_innovation_stability_gate.json")
+ args = parser.parse_args()
+ records = {}
+ source_commits = set()
+ for label, margin in MARGINS:
+ path = os.path.join(args.input_dir, f"margin_{label}.json")
+ with open(path) as handle:
+ record = json.load(handle)
+ records[label] = record
+ expected = {
+ "protocol": "kp_mixed_traffic_nonfinite_diagnosis_v1",
+ "scope": "training_only_no_validation_or_test_evaluation",
+ "rule": "innovation",
+ "predictor_mode": "closed_form",
+ "predictor_every": 0,
+ "max_steps": 352,
+ "stability_margin": margin,
+ "validation_evaluations": 0,
+ "test_evaluations": 0,
+ }
+ for key, value in expected.items():
+ if record.get(key) != value:
+ raise ValueError(f"S0 {label} {key} drift")
+ if record["provenance"]["git_tracked_dirty"]:
+ raise ValueError(f"tracked-dirty S0 {label}")
+ source_commits.add(record["provenance"]["git_commit"])
+ if record["split"]["validation_index_sha256"] != SPLIT_HASH:
+ raise ValueError(f"S0 {label} split drift")
+ if len(source_commits) != 1:
+ raise ValueError("S0 candidates must share one source revision")
+
+ metrics = {}
+ eligible = []
+ checks = {}
+ for label, margin in MARGINS:
+ record = records[label]
+ trajectory = record["trajectory"]
+ if len(trajectory) != 352:
+ raise ValueError(f"S0 {label} trajectory is incomplete")
+ all_finite = all(math.isfinite(value) for value in
+ numeric_leaves(record))
+ state_finite = all(
+ value["all_finite"]
+ for row in trajectory
+ for value in row["parameter_state"].values())
+ losses = [float(row["batch_loss"]) for row in trajectory]
+ ratios = [float(row["teaching_rms"])
+ / max(float(row["instruction_rms"]), 1e-30)
+ for row in trajectory]
+ forward_weight_max = max(float(
+ row["parameter_state"]["forward_weight"][
+ "max_abs_over_finite_tensors"]) for row in trajectory)
+ feedback_weight_max = max(float(
+ row["parameter_state"]["feedback_weight"][
+ "max_abs_over_finite_tensors"]) for row in trajectory)
+ forward_momentum_max = max(float(
+ row["parameter_state"]["forward_momentum"][
+ "max_abs_over_finite_tensors"]) for row in trajectory)
+ feedback_momentum_max = max(float(
+ row["parameter_state"]["feedback_momentum"][
+ "max_abs_over_finite_tensors"]) for row in trajectory)
+ fit = record["predictor_warmup"]["closed_form_fit"]
+ ratio_errors = [abs(float(value) - 4.0) for value in
+ record["traffic_calibration"][
+ "realized_traffic_instruction_rms_ratio"]]
+ candidate_checks = {
+ "record_trajectory_and_state_finite": (
+ all_finite and state_finite
+ and record["first_any_nonfinite_step"] is None
+ and record["first_training_failure_step"] is None),
+ "maximum_batch_loss_at_most_10": max(losses) <= 10.0,
+ "final_32_mean_loss_at_most_2p5": (
+ statistics.mean(losses[-32:]) <= 2.5),
+ "maximum_used_instruction_rms_ratio_at_most_2": (
+ max(ratios) <= 2.0),
+ "forward_and_feedback_weight_max_at_most_10": (
+ forward_weight_max <= 10.0 and feedback_weight_max <= 10.0),
+ "forward_and_feedback_momentum_max_at_most_50": (
+ forward_momentum_max <= 50.0
+ and feedback_momentum_max <= 50.0),
+ "one_sided_closed_form_certificate": (
+ int(fit["observations"]) == 64
+ and float(fit["max_positive_residual_soma_slope"]) <= 1e-7),
+ "predictor_frozen_during_task": all(
+ row["predictor_updated"] is False for row in trajectory),
+ "traffic_ratio_calibrated": max(ratio_errors) <= 1e-5,
+ "task_loader_state_restored": (
+ record["predictor_warmup"][
+ "task_loader_state_restored"] is True),
+ "no_held_out_evaluations": (
+ record["validation_evaluations"] == 0
+ and record["test_evaluations"] == 0),
+ }
+ passed = all(candidate_checks.values())
+ if passed:
+ eligible.append((margin, label))
+ checks[label] = candidate_checks
+ metrics[label] = {
+ "margin": margin,
+ "eligible": passed,
+ "maximum_batch_loss": max(losses),
+ "final_32_mean_loss": statistics.mean(losses[-32:]),
+ "maximum_used_instruction_rms_ratio": max(ratios),
+ "maximum_forward_weight": forward_weight_max,
+ "maximum_feedback_weight": feedback_weight_max,
+ "maximum_forward_momentum": forward_momentum_max,
+ "maximum_feedback_momentum": feedback_momentum_max,
+ "max_positive_residual_soma_slope": float(
+ fit["max_positive_residual_soma_slope"]),
+ "min_residual_soma_slope": float(
+ fit["min_residual_soma_slope"]),
+ "maximum_initial_traffic_ratio_error": max(ratio_errors),
+ }
+ selected = min(eligible)[1] if eligible else None
+ status = "passed" if selected is not None else "failed_no_eligible_margin"
+ output = {
+ "protocol": "kp_stability_margin_training_prefix_v1",
+ "status": status,
+ "selected": selected,
+ "checks": checks,
+ "metrics": metrics,
+ "source_commit": next(iter(source_commits)),
+ "validation_evaluations": 0,
+ "test_evaluations": 0,
+ "review_score_before": 5,
+ "review_score_after": 5,
+ "score_change_rule": (
+ "training-only stability selection cannot change the paper 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(output, indent=2))
+
+
+if __name__ == "__main__":
+ main()
+