#!/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") # Once a sham control has become nonfinite, numerical orthogonality of a # projection that is not applied to its update is undefined. The frozen # protocol explicitly retains such terminal collapses. Observation count # and instruction leakage remain auditable for every batch; numerical # projection quality is required only for a finite trajectory. Pretask # predictor quality remains required for both cases. require(finite(warmup["post_warmup_traffic_residual_rms_ratio"]) and warmup["post_warmup_traffic_residual_rms_ratio"] <= 1e-5, f"{path}: pretask predictor residual") if row["final"]["finite"]: 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") analysis_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") record_paths = { rule: os.path.join(args.outdir, f"{rule}.json") for rule in RULES } record_commits = set() for path in record_paths.values(): with open(path) as handle: record_commits.add(json.load(handle)["provenance"]["git_commit"]) require(len(record_commits) == 1, "raw and matched controls must share one training source") training_commit = record_commits.pop() critical_training_paths = [ "KP_TEACHING_SIGNAL_ABLATION.md", "experiments/conv_local_smoke.py", "experiments/conv_run.py", "experiments/kp_teaching_signal_ablation.py", "sdil/conv.py", ] require(subprocess.run( ["git", "diff", "--quiet", training_commit, analysis_commit, "--", *critical_training_paths], cwd=ROOT).returncode == 0, "analysis revision changes a frozen training-critical path") rows = { rule: validate_record( record_paths[rule], rule, training_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()), "training_source_commit": training_commit, "analysis_source_commit": analysis_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()