#!/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()