diff options
Diffstat (limited to 'experiments/analyze_kp_innovation_full.py')
| -rw-r--r-- | experiments/analyze_kp_innovation_full.py | 194 |
1 files changed, 194 insertions, 0 deletions
diff --git a/experiments/analyze_kp_innovation_full.py b/experiments/analyze_kp_innovation_full.py new file mode 100644 index 0000000..2b5394d --- /dev/null +++ b/experiments/analyze_kp_innovation_full.py @@ -0,0 +1,194 @@ +#!/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 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") + 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"]), + ]) + + all_finite = all(record["final"]["finite"] for record in records.values()) + all_finite = all_finite and all(math.isfinite(value) for value in ( + trajectory_values + list(accuracies.values()) + [innovation_early, + innovation_raw_early])) + 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 + + 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() |
