#!/usr/bin/env python3 """Audit and gate the frozen MT-2 full mixed-traffic validation panel.""" import argparse import json import math import os RULES = ("raw", "matched", "innovation") SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b" def mean_early(values): count = max(1, len(values) // 3) return sum(float(value) for value in values[:count]) / count def numeric_leaves(value): """Yield every numeric audit value while excluding boolean flags.""" 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 require_pass(path, protocol): with open(path) as handle: record = json.load(handle) if record.get("protocol") != protocol or record.get("status") != "passed": raise ValueError(f"{protocol} did not pass") return record def main(): parser = argparse.ArgumentParser() parser.add_argument("--input_dir", default="results/kp_innovation_full") parser.add_argument( "--short_gate", default="results/kp_innovation_short_gate.json") parser.add_argument("--kp_full_gate", default="results/kp_full_gate.json") parser.add_argument("--out", default="results/kp_innovation_full_gate.json") args = parser.parse_args() require_pass(args.short_gate, "kp_mixed_traffic_short_v1") kp_gate = require_pass(args.kp_full_gate, "kolen_pollack_full_v1") kp_accuracy = float(kp_gate["metrics"]["accuracy"]) bp_macs = int(kp_gate["metrics"]["bp_total_macs"]) expected_common = { "mode": "kp_traffic", "depth": 20, "width": 16, "seed": 0, "loader_seed": 0, "batch_size": 128, "epochs": 200, "train_limit": 0, "val_examples": 5000, "split_seed": 2027, "eval_split": "validation", "eval_every": 0, "augment_train": 1, "lr": 0.1, "output_lr": 0.1, "lr_schedule": "step", "lr_milestones": "100,150", "lr_gamma": 0.1, "warmup_epochs": 0, "momentum": 0.9, "weight_decay": 1e-4, "normalization": "batchnorm", "a_scale": 1.0, "traffic_seed": 4000, "traffic_ratio": 4.0, "traffic_calibration_examples": 64, "learn_P": 1, "eta_P": 0.1, "predictor_warmup_steps": 20, "predictor_every": 16, "alignment_probe": 32, } records = {} source_commits = set() for rule in RULES: with open(os.path.join(args.input_dir, f"{rule}.json")) as handle: record = json.load(handle) records[rule] = record for key, value in {**expected_common, "traffic_rule": rule}.items(): if record["args"].get(key) != value: raise ValueError(f"MT-2 {rule} {key} drift") if record["provenance"]["git_tracked_dirty"]: raise ValueError(f"tracked-dirty MT-2 {rule} result") source_commits.add(record["provenance"]["git_commit"]) if record["split"]["validation_index_sha256"] != SPLIT_HASH: raise ValueError(f"MT-2 {rule} split drift") protocol = record["evaluation_protocol"] if protocol["test_evaluations"] or protocol["test_used_for_selection"]: raise ValueError(f"MT-2 {rule} touched test") if record.get("calibration_metric_space") != ( "reciprocal_local_activity_products_with_mixed_apical_traffic"): raise ValueError(f"MT-2 {rule} metric-space drift") warmup = record.get("predictor_warmup", {}) if (warmup.get("instruction_present") is not False or warmup.get("task_loader_state_restored") is not True): raise ValueError(f"MT-2 {rule} neutral-warmup invariant failed") if len(source_commits) != 1: raise ValueError("MT-2 conditions must share one source revision") accuracies = {rule: float(records[rule]["final"]["accuracy"]) for rule in RULES} diagnostics = {rule: records[rule]["diagnostics"] for rule in RULES} innovation_early = float(diagnostics["innovation"]["early_third_mean"]) innovation_raw_early = mean_early( diagnostics["innovation"]["raw_negative_gradient_cosine"]) tracking = { rule: [row.get("feedback_tracking") for row in records[rule]["epochs"]] for rule in RULES } trajectory_values = [] for rule in RULES: if len(tracking[rule]) != 200 or any(value is None for value in tracking[rule]): raise ValueError(f"MT-2 {rule} tracking trajectory is incomplete") for row, values in zip(records[rule]["epochs"], tracking[rule]): mixed = row.get("mixed_apical") if mixed is None: raise ValueError(f"MT-2 {rule} mixed-apical trajectory is incomplete") trajectory_values.extend([ float(row["train_loss"]), float(values["mean_feedback_forward_cosine"]), float(values["mean_feedback_forward_relative_error"]), float(values["min_feedback_forward_cosine"]), float(values["max_feedback_forward_relative_error"]), float(mixed["teaching_rms"]), float(mixed["instruction_rms"]), float(mixed["raw_apical_rms"]), float(mixed["innovation_rms"]), float(mixed["traffic_rms"]), ]) initial_ratio_errors = [] total_macs = {} queries = {} elementwise = {} for rule, record in records.items(): initial_ratio_errors.extend(abs(float(value) - 4.0) for value in record["traffic_calibration"]["realized_traffic_instruction_rms_ratio"]) total_macs[rule] = int(record["work"]["total_macs_estimate"]) queries[rule] = int(record["work"]["logical_batch_loss_queries"]) elementwise[rule] = int(record["work"]["elementwise_operations_estimate"]) matched_norm_error = float( diagnostics["matched"]["max_norm_match_relative_error"]) predictor_residual_ratios = { rule: float(value["predictor_traffic_residual_rms_ratio"]) for rule, value in diagnostics.items() } final_feedback_cosine = float( diagnostics["innovation"]["mean_feedback_forward_cosine"]) late_feedback_cosine = sum(float(value["mean_feedback_forward_cosine"]) for value in tracking["innovation"][150:]) / 50 audit_values = trajectory_values + [innovation_early, innovation_raw_early, final_feedback_cosine, late_feedback_cosine, matched_norm_error] for record in records.values(): audit_values.extend(numeric_leaves(record["final"])) audit_values.extend(numeric_leaves(record["diagnostics"])) audit_values.extend(numeric_leaves(record["traffic_calibration"])) audit_values.extend(numeric_leaves(record["predictor_warmup"])) all_finite = all(record["final"]["finite"] for record in records.values()) all_finite = all_finite and all( math.isfinite(value) for value in audit_values) checks = { "records_trajectories_and_diagnostics_finite": all_finite, "innovation_accuracy_at_least_0.88": accuracies["innovation"] >= 0.88, "innovation_within_3_points_of_clean_kp": ( accuracies["innovation"] >= kp_accuracy - 0.03), "innovation_gain_over_raw_at_least_0.05": ( accuracies["innovation"] - accuracies["raw"] >= 0.05), "innovation_gain_over_matched_at_least_0.03": ( accuracies["innovation"] - accuracies["matched"] >= 0.03), "innovation_early_alignment_at_least_0.80": innovation_early >= 0.80, "same_network_alignment_gain_over_raw_at_least_0.15": ( innovation_early - innovation_raw_early >= 0.15), "final_feedback_cosine_at_least_0.95": final_feedback_cosine >= 0.95, "epoch151_to200_feedback_cosine_at_least_0.95": ( late_feedback_cosine >= 0.95), "initial_layer_ratio_error_at_most_1e-5": ( max(initial_ratio_errors) <= 1e-5), "post_training_predictor_residual_ratio_at_most_0.05": ( max(predictor_residual_ratios.values()) <= 0.05), "matched_norm_relative_error_at_most_1e-6": matched_norm_error <= 1e-6, "zero_task_loss_queries": all(value == 0 for value in queries.values()), "each_affine_mac_count_at_most_1.40x_bp": all( value <= 1.40 * bp_macs for value in total_macs.values()), "elementwise_cost_reported": all(value > 0 for value in elementwise.values()), } status = "passed" if all(checks.values()) else "failed" output = { "protocol": "kp_mixed_traffic_full_v1", "status": status, "checks": checks, "metrics": { "accuracy": accuracies, "clean_kp_accuracy": kp_accuracy, "innovation_gain_over_raw": ( accuracies["innovation"] - accuracies["raw"]), "innovation_gain_over_matched": ( accuracies["innovation"] - accuracies["matched"]), "innovation_early_third_alignment": innovation_early, "same_innovation_network_raw_early_third_alignment": innovation_raw_early, "final_feedback_forward_cosine": final_feedback_cosine, "epoch151_to200_feedback_forward_cosine": late_feedback_cosine, "max_initial_layer_ratio_error": max(initial_ratio_errors), "post_training_predictor_residual_rms_ratio": ( predictor_residual_ratios), "matched_norm_relative_error": matched_norm_error, "total_macs": total_macs, "mac_ratio_to_bp": {rule: value / bp_macs for rule, value in total_macs.items()}, "elementwise_operations_estimate": elementwise, "logical_batch_loss_queries": queries, "source_commit": next(iter(source_commits)), }, "test_confirmation_opened": status == "passed", "confirmation_test_seeds_touched": False, "review_score_before": 5, "review_score_after": 5, "score_change_rule": ( "a one-seed full validation result cannot establish acceptance"), } 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()