summaryrefslogtreecommitdiff
path: root/experiments/analyze_kp_teaching_signal.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-26 12:27:52 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-26 12:27:52 -0500
commit00baa870b8b61a0974438233d9240eebb40c384f (patch)
treebae281f38fe7512a0e0ad46da43f6262eb3c94dc /experiments/analyze_kp_teaching_signal.py
parent8f9364e8ee2e2a19eb83b58452fa0e8cb155353c (diff)
experiment: freeze KP teaching-signal ablation
Diffstat (limited to 'experiments/analyze_kp_teaching_signal.py')
-rw-r--r--experiments/analyze_kp_teaching_signal.py275
1 files changed, 275 insertions, 0 deletions
diff --git a/experiments/analyze_kp_teaching_signal.py b/experiments/analyze_kp_teaching_signal.py
new file mode 100644
index 0000000..db8649a
--- /dev/null
+++ b/experiments/analyze_kp_teaching_signal.py
@@ -0,0 +1,275 @@
+#!/usr/bin/env python3
+"""Audit the frozen KTS-1 raw-KP teaching-signal controls."""
+import argparse
+import hashlib
+import json
+import math
+import os
+import statistics
+import subprocess
+
+from kp_teaching_signal_ablation import ROOT, RULES, command
+
+
+def require(condition, message):
+ if not condition:
+ raise ValueError(message)
+
+
+def sha256(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
+ digest.update(chunk)
+ return digest.hexdigest()
+
+
+def close(left, right, tolerance=1e-12):
+ return abs(float(left) - float(right)) <= tolerance
+
+
+def finite(value):
+ return isinstance(value, (int, float)) and math.isfinite(value)
+
+
+def same_scalar(left, right, tolerance=1e-7):
+ left = float(left)
+ right = float(right)
+ if math.isnan(left) or math.isnan(right):
+ return math.isnan(left) and math.isnan(right)
+ if math.isinf(left) or math.isinf(right):
+ return left == right
+ return abs(left - right) <= tolerance
+
+
+def same_list(left, right, tolerance=1e-7):
+ return len(left) == len(right) and all(
+ same_scalar(a, b, tolerance) for a, b in zip(left, right))
+
+
+def validate_record(path, rule, source_commit, bp_reference_macs):
+ with open(path) as handle:
+ row = json.load(handle)
+ require(row["provenance"]["git_commit"] == source_commit,
+ f"{path}: source commit")
+ require(row["provenance"]["git_tracked_dirty"] is False,
+ f"{path}: dirty source")
+ expected = command("cuda", rule, path)[2:-2]
+ expected_args = {}
+ for index in range(0, len(expected), 2):
+ key = expected[index].removeprefix("--")
+ expected_args[key] = expected[index + 1]
+ actual = row["args"]
+ numeric = {
+ "depth": int, "width": int, "seed": int, "loader_seed": int,
+ "batch_size": int, "epochs": int, "train_limit": int,
+ "val_examples": int, "split_seed": int, "eval_every": int,
+ "augment_train": int, "warmup_epochs": int, "learn_P": int,
+ "predictor_warmup_steps": int, "predictor_every": int,
+ "neutral_projection": int, "traffic_seed": int,
+ "traffic_calibration_examples": int, "alignment_probe": int,
+ "lr_gamma": float, "momentum": float, "weight_decay": float,
+ "lr": float, "output_lr": float, "eta_P": float,
+ "traffic_ratio": float,
+ }
+ for key, expected_value in expected_args.items():
+ if key in ("device", "out"):
+ continue
+ caster = numeric.get(key, str)
+ require(actual[key] == caster(expected_value),
+ f"{path}: argument {key}")
+
+ require(row["split"]["train_examples"] == 45000, f"{path}: train split")
+ require(row["split"]["validation_examples"] == 5000,
+ f"{path}: validation split")
+ require(row["split"]["test_examples"] == 10000, f"{path}: test split")
+ require(row["evaluation_protocol"] == {
+ "validation_evaluations": 1,
+ "test_evaluations": 0,
+ "test_used_for_selection": False,
+ }, f"{path}: evaluation protocol")
+ counters = row["counters"]
+ require(counters["ordinary_examples"] == 9_000_000,
+ f"{path}: ordinary observations")
+ require(counters["neutral_projection_examples"]
+ == counters["ordinary_examples"],
+ f"{path}: projection observation budget")
+ require(counters["predictor_update_examples"] == 64,
+ f"{path}: slow fit observations")
+ require(counters["logical_batch_loss_queries"] == 0,
+ f"{path}: task loss queries")
+ require(counters["causal_scalar_observations"] == 0,
+ f"{path}: causal observations")
+ require(row["work"]["neutral_projection_observations"]
+ == counters["ordinary_examples"],
+ f"{path}: reported projection observations")
+ require(row["work"]["elementwise_operations_estimate"] > 0,
+ f"{path}: elementwise work")
+ require(row["work"]["total_macs_estimate"]
+ <= 1.34 * bp_reference_macs,
+ f"{path}: MAC ceiling")
+
+ calibration = row["traffic_calibration"]
+ require(max(abs(float(value) - 4.0) for value in
+ calibration["realized_traffic_instruction_rms_ratio"])
+ <= 1e-5, f"{path}: traffic ratio")
+ warmup = row["predictor_warmup"]
+ require(warmup["mode"] == "closed_form" and warmup["examples"] == 64,
+ f"{path}: closed-form fit")
+ require(warmup["closed_form_fit"]["stability_margin"] == 0.0,
+ f"{path}: predictor margin")
+
+ projections = [epoch["neutral_projection"] for epoch in row["epochs"]]
+ require(all(value["instruction_observations"] == 0
+ for value in projections),
+ f"{path}: projection instruction leakage")
+ require(min(value["minimum_observations"] for value in projections) > 0
+ and max(value["maximum_observations"]
+ for value in projections) <= 128,
+ f"{path}: projection batch observations")
+ require(max(value["maximum_post_projection_traffic_rms_ratio"]
+ for value in projections) <= 1e-5,
+ f"{path}: projected traffic remainder")
+ require(max(value["maximum_absolute_post_projection_soma_slope"]
+ for value in projections) <= 1e-5,
+ f"{path}: projected soma slope")
+
+ diagnostics = row["diagnostics"]
+ if rule == "raw":
+ require(same_list(
+ diagnostics["used_negative_gradient_cosine"],
+ diagnostics["raw_negative_gradient_cosine"]),
+ f"{path}: raw direction")
+ else:
+ require(same_list(
+ diagnostics["used_negative_gradient_cosine"],
+ diagnostics["matched_negative_gradient_cosine"]),
+ f"{path}: matched direction")
+ if row["final"]["finite"]:
+ require(diagnostics["max_norm_match_relative_error"] <= 1e-6,
+ f"{path}: matched norm")
+ require(diagnostics["max_norm_match_direction_error"] <= 1e-6,
+ f"{path}: matched direction preservation")
+
+ for epoch in row["epochs"]:
+ mixed = epoch["mixed_apical"]
+ if all(finite(value) for value in mixed.values()):
+ if rule == "raw":
+ require(close(
+ mixed["teaching_rms"], mixed["raw_apical_rms"], 1e-10),
+ f"{path}: epoch raw signal")
+ else:
+ require(close(
+ mixed["teaching_rms"], mixed["innovation_rms"], 1e-9),
+ f"{path}: epoch matched norm")
+ require(epoch["neutral_projection"]["instruction_observations"] == 0,
+ f"{path}: epoch projection leakage")
+ accuracy = row["final"]["accuracy"]
+ require(finite(accuracy) and 0.0 <= accuracy <= 1.0,
+ f"{path}: endpoint accuracy")
+ return row
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--outdir", default="results/kp_teaching_signal_ablation")
+ parser.add_argument(
+ "--dynamic",
+ default="results/kp_dynamic_projection_full/dynamic.json")
+ parser.add_argument(
+ "--prerequisite",
+ default="results/kp_dynamic_projection_full_gate.json")
+ parser.add_argument(
+ "--out", default="results/kp_teaching_signal_ablation_gate.json")
+ args = parser.parse_args()
+
+ tracked_dirty = subprocess.run(
+ ["git", "status", "--porcelain", "--untracked-files=no"],
+ cwd=ROOT, check=True, capture_output=True, text=True).stdout.strip()
+ require(not tracked_dirty, "analysis requires a clean tracked source")
+ source_commit = subprocess.run(
+ ["git", "rev-parse", "HEAD"], cwd=ROOT, check=True,
+ capture_output=True, text=True).stdout.strip()
+ with open(args.prerequisite) as handle:
+ prerequisite = json.load(handle)
+ require(prerequisite["protocol"]
+ == "kp_dynamic_neutral_projection_full_v1",
+ "wrong prerequisite protocol")
+ require(prerequisite["status"] == "passed", "prerequisite did not pass")
+ with open(args.dynamic) as handle:
+ dynamic = json.load(handle)
+ require(close(dynamic["final"]["accuracy"],
+ prerequisite["metrics"]["accuracy"]),
+ "dynamic record/gate mismatch")
+
+ rows = {
+ rule: validate_record(
+ os.path.join(args.outdir, f"{rule}.json"), rule, source_commit,
+ prerequisite["metrics"]["bp_total_macs"])
+ for rule in RULES
+ }
+ dynamic_accuracy = float(dynamic["final"]["accuracy"])
+ raw_accuracy = float(rows["raw"]["final"]["accuracy"])
+ matched_accuracy = float(rows["matched"]["final"]["accuracy"])
+ raw_gap = dynamic_accuracy - raw_accuracy
+ matched_gap = dynamic_accuracy - matched_accuracy
+ raw_finite = bool(rows["raw"]["final"]["finite"])
+ matched_finite = bool(rows["matched"]["final"]["finite"])
+ checks = {
+ "two_frozen_controls_present": len(rows) == 2,
+ "dynamic_reference_is_finite": bool(dynamic["final"]["finite"]),
+ "raw_degrades_or_collapses_by_at_least_5_points": (
+ (not raw_finite) or raw_gap >= 0.05),
+ "matched_degrades_or_collapses_by_at_least_3_points": (
+ (not matched_finite) or matched_gap >= 0.03),
+ "raw_and_matched_use_identical_source": (
+ len({row["provenance"]["git_commit"]
+ for row in rows.values()}) == 1),
+ "all_control_invariants_audited": True,
+ }
+ status = "passed" if all(checks.values()) else "failed"
+ report = {
+ "protocol": "kp_teaching_signal_ablation_validation_v1",
+ "status": status,
+ "checks": checks,
+ "metrics": {
+ "dynamic_accuracy": dynamic_accuracy,
+ "clean_kp_accuracy": prerequisite["metrics"][
+ "clean_kp_accuracy"],
+ "raw_accuracy": raw_accuracy,
+ "matched_accuracy": matched_accuracy,
+ "dynamic_minus_raw_points": 100.0 * raw_gap,
+ "dynamic_minus_matched_points": 100.0 * matched_gap,
+ "raw_finite": raw_finite,
+ "matched_finite": matched_finite,
+ "control_mean_wall_s": statistics.mean(
+ row["timing"]["total_timed_wall_s"]
+ for row in rows.values()),
+ "source_commit": source_commit,
+ "prerequisite_sha256": sha256(args.prerequisite),
+ },
+ "claim": (
+ "Under four-RMS soma-predictable apical traffic, innovation "
+ "subtraction is necessary relative to both raw and norm-matched "
+ "KP controls."
+ if status == "passed" else
+ "The frozen validation endpoint does not establish a KP "
+ "teaching-signal advantage."
+ ),
+ "scope_limit": (
+ "This is a conditional robustness result on one validation seed; "
+ "it is not clean-setting or strong-baseline superiority."
+ ),
+ }
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ with open(args.out, "w") as handle:
+ json.dump(report, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps(report, indent=2, sort_keys=True))
+ if status != "passed":
+ raise SystemExit(1)
+
+
+if __name__ == "__main__":
+ main()