summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py219
1 files changed, 219 insertions, 0 deletions
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()