From 00baa870b8b61a0974438233d9240eebb40c384f Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Sun, 26 Jul 2026 12:27:52 -0500 Subject: experiment: freeze KP teaching-signal ablation --- experiments/analyze_kp_teaching_signal.py | 275 ++++++++++++++++++++++++++++++ 1 file changed, 275 insertions(+) create mode 100644 experiments/analyze_kp_teaching_signal.py (limited to 'experiments/analyze_kp_teaching_signal.py') 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() -- cgit v1.2.3