From efe356b3d9b99aa397caf8e8d70320dedc8d4450 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 16:32:38 -0500 Subject: diagnostic: localize mixed-traffic nonfinite onset --- experiments/diagnose_kp_traffic_nonfinite.py | 219 +++++++++++++++++++++++++++ 1 file changed, 219 insertions(+) create mode 100644 experiments/diagnose_kp_traffic_nonfinite.py diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py new file mode 100644 index 0000000..02d18ae --- /dev/null +++ b/experiments/diagnose_kp_traffic_nonfinite.py @@ -0,0 +1,219 @@ +#!/usr/bin/env python3 +"""Training-only localization of the frozen MT-1 nonfinite failure. + +This reproduces MT-1 initialization, traffic calibration, neutral predictor +warmup, and task minibatch order, but performs no validation or test +evaluation. It stops after the first update that makes any model, optimizer, +feedback, predictor, or BatchNorm tensor nonfinite. +""" +import argparse +import json +import math +import os +import subprocess +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.conv import (CIFARKPMixedTrafficResNet, ConvSDILConfig, + conv_kp_mixed_traffic_step) +from sdil.data import DATA_DIR, get_cifar_image_splits + + +ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +def provenance(): + def run(command): + return subprocess.run( + command, cwd=ROOT, check=True, capture_output=True, + text=True).stdout.strip() + return { + "git_commit": run(["git", "rev-parse", "HEAD"]), + "git_tracked_dirty": bool(run( + ["git", "status", "--porcelain", "--untracked-files=no"])), + } + + +def tensor_groups(net): + return { + "forward_weight": net.W, + "forward_bn_scale": net.gamma, + "forward_bn_bias": net.beta, + "forward_readout": [net.W_out, net.b_out], + "feedback_weight": net.Q, + "feedback_readout": [net.R_out], + "predictor_slope": net.P_traffic, + "predictor_bias": net.P_traffic_bias, + "forward_momentum": net.mW, + "forward_bn_scale_momentum": net.mgamma, + "forward_bn_bias_momentum": net.mbeta, + "forward_readout_momentum": [net.mW_out, net.mb_out], + "feedback_momentum": net.mQ, + "feedback_readout_momentum": [net.mR_out], + "bn_running_mean": net.running_mean, + "bn_running_var": net.running_var, + } + + +def network_state(net): + output = {} + for name, tensors in tensor_groups(net).items(): + bad = [] + max_abs = 0.0 + for index, value in enumerate(tensors): + finite = bool(torch.isfinite(value).all()) + if finite: + max_abs = max(max_abs, float(value.abs().max())) + else: + bad.append(index) + output[name] = { + "all_finite": not bad, + "nonfinite_indices": bad, + "max_abs_over_finite_tensors": max_abs, + } + return output + + +def nonfinite_groups(state): + return {name: value["nonfinite_indices"] + for name, value in state.items() if not value["all_finite"]} + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--rule", choices=("raw", "matched", "innovation"), + default="innovation") + parser.add_argument("--device", default="cuda") + parser.add_argument("--max_steps", type=int, default=352) + parser.add_argument("--out", required=True) + args = parser.parse_args() + if args.max_steps < 1: + raise ValueError("max_steps must be positive") + + torch.manual_seed(0) + if str(args.device).startswith("cuda"): + if not torch.cuda.is_available(): + raise RuntimeError("CUDA device requested but unavailable") + torch.cuda.manual_seed_all(0) + train, _, _, input_shape, n_out, split = get_cifar_image_splits( + batch_size=128, data_dir=DATA_DIR, device=args.device, + val_examples=5000, split_seed=2027, loader_seed=0, + augment_train=True) + if input_shape != (3, 32, 32) or n_out != 10: + raise AssertionError("unexpected CIFAR dimensions") + net = CIFARKPMixedTrafficResNet( + depth=20, base_width=16, n_classes=10, device=args.device, seed=0, + weight_scale=1.0, residual_scale=1.0, normalization="batchnorm", + feedback_seed=None, feedback_scale=1.0, traffic_seed=4000) + net.traffic_rule = args.rule + config = ConvSDILConfig( + eta=0.1, eta_output=0.1, eta_A=0.0, eta_P=0.1, + momentum=0.9, weight_decay=1e-4, learn_A=False, learn_P=True) + config.validate() + + loader_state = train.g.get_state().clone() + calibration_x = train.x[:64] + calibration_y = train.y[:64] + calibration_forward = net.forward( + calibration_x, return_cache=True, training=True, update_stats=False) + calibration_error = ( + torch.softmax(calibration_forward["logits"], dim=1) + - torch.nn.functional.one_hot( + calibration_y, 10).to(calibration_forward["logits"].dtype)) + calibration_instruction = net.hierarchical_teaching( + calibration_error, calibration_forward) + calibration = net.calibrate_traffic_gain( + calibration_instruction, calibration_forward["hiddens"], 4.0) + del calibration_forward, calibration_error, calibration_instruction + + iterator = iter(train) + warmup_mse = [] + for _ in range(20): + x, _ = next(iterator) + forward = net.forward(x, training=True, update_stats=False) + warmup_mse.append(net.predictor_step(forward["hiddens"], 0.1)) + train.g.set_state(loader_state) + loader_state_restored = torch.equal(train.g.get_state(), loader_state) + audit_forward = net.forward( + calibration_x, training=True, update_stats=False) + post_warmup_ratio = net.predictor_traffic_residual_rms_ratio( + audit_forward["hiddens"]) + del forward, audit_forward + + initial_state = network_state(net) + if nonfinite_groups(initial_state): + raise AssertionError("network is nonfinite before task training") + trajectory = [] + first_nonfinite_step = None + first_nonfinite = {} + for step, (x, y) in enumerate(train): + if step >= args.max_steps: + break + result = conv_kp_mixed_traffic_step( + net, x, y, config, step, args.rule, predictor_every=16) + state = network_state(net) + bad = nonfinite_groups(state) + trajectory.append({ + "step": step + 1, + "batch_loss": result["loss"], + "teaching_rms": result["teaching_rms"], + "instruction_rms": result["instruction_rms"], + "raw_apical_rms": result["raw_apical_rms"], + "innovation_rms": result["innovation_rms"], + "traffic_rms": result["traffic_rms"], + "predictor_updated": result["did_predictor_update"], + "predictor_mse": result["predictor_mse"], + "parameter_state": state, + }) + if bad: + first_nonfinite_step = step + 1 + first_nonfinite = bad + break + + numeric = [] + for row in trajectory: + numeric.extend(float(row[key]) for key in ( + "batch_loss", "teaching_rms", "instruction_rms", + "raw_apical_rms", "innovation_rms", "traffic_rms")) + if row["predictor_mse"] is not None: + numeric.append(float(row["predictor_mse"])) + output = { + "protocol": "kp_mixed_traffic_nonfinite_diagnosis_v1", + "scope": "training_only_no_validation_or_test_evaluation", + "provenance": provenance(), + "rule": args.rule, + "max_steps": args.max_steps, + "split": split, + "traffic_calibration": calibration, + "predictor_warmup": { + "steps": 20, + "first_mse": warmup_mse[0], + "last_mse": warmup_mse[-1], + "post_warmup_traffic_residual_rms_ratio": post_warmup_ratio, + "task_loader_state_restored": loader_state_restored, + }, + "initial_parameter_state": initial_state, + "trajectory": trajectory, + "first_nonfinite_step": first_nonfinite_step, + "first_nonfinite_groups": first_nonfinite, + "trajectory_metrics_finite": all(math.isfinite(value) + for value in numeric), + "validation_evaluations": 0, + "test_evaluations": 0, + } + 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({ + "out": args.out, + "first_nonfinite_step": first_nonfinite_step, + "first_nonfinite_groups": first_nonfinite, + "steps_recorded": len(trajectory), + }, indent=2)) + + +if __name__ == "__main__": + main() -- cgit v1.2.3