summaryrefslogtreecommitdiff
path: root/experiments/diagnose_kp_traffic_nonfinite.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:33:44 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:33:44 -0500
commitd2f7f35f112e5e21515c134548d6d27ada5a0b4f (patch)
tree79813f49ca7e4df5fb6261785204971b5d9f3fba /experiments/diagnose_kp_traffic_nonfinite.py
parentefe356b3d9b99aa397caf8e8d70320dedc8d4450 (diff)
diagnostic: separate training-active nonfinite state
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py47
1 files changed, 38 insertions, 9 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py
index 02d18ae..c144062 100644
--- a/experiments/diagnose_kp_traffic_nonfinite.py
+++ b/experiments/diagnose_kp_traffic_nonfinite.py
@@ -22,6 +22,7 @@ from sdil.data import DATA_DIR, get_cifar_image_splits
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+TRAINING_INACTIVE_GROUPS = {"bn_running_mean", "bn_running_var"}
def provenance():
@@ -146,8 +147,12 @@ def main():
if nonfinite_groups(initial_state):
raise AssertionError("network is nonfinite before task training")
trajectory = []
- first_nonfinite_step = None
- first_nonfinite = {}
+ first_any_nonfinite_step = None
+ first_any_nonfinite = {}
+ first_nonfinite_by_group = {}
+ first_training_failure_step = None
+ first_training_failure_groups = {}
+ first_nonfinite_metrics = []
for step, (x, y) in enumerate(train):
if step >= args.max_steps:
break
@@ -155,6 +160,22 @@ def main():
net, x, y, config, step, args.rule, predictor_every=16)
state = network_state(net)
bad = nonfinite_groups(state)
+ for name, indices in bad.items():
+ first_nonfinite_by_group.setdefault(name, {
+ "step": step + 1, "indices": indices})
+ if bad and first_any_nonfinite_step is None:
+ first_any_nonfinite_step = step + 1
+ first_any_nonfinite = bad
+ metric_names = [
+ key for key in (
+ "loss", "teaching_rms", "instruction_rms",
+ "raw_apical_rms", "innovation_rms", "traffic_rms")
+ if not math.isfinite(float(result[key]))]
+ if (result["predictor_mse"] is not None
+ and not math.isfinite(float(result["predictor_mse"]))):
+ metric_names.append("predictor_mse")
+ active_bad = {name: indices for name, indices in bad.items()
+ if name not in TRAINING_INACTIVE_GROUPS}
trajectory.append({
"step": step + 1,
"batch_loss": result["loss"],
@@ -167,9 +188,10 @@ def main():
"predictor_mse": result["predictor_mse"],
"parameter_state": state,
})
- if bad:
- first_nonfinite_step = step + 1
- first_nonfinite = bad
+ if active_bad or metric_names:
+ first_training_failure_step = step + 1
+ first_training_failure_groups = active_bad
+ first_nonfinite_metrics = metric_names
break
numeric = []
@@ -196,8 +218,12 @@ def main():
},
"initial_parameter_state": initial_state,
"trajectory": trajectory,
- "first_nonfinite_step": first_nonfinite_step,
- "first_nonfinite_groups": first_nonfinite,
+ "first_any_nonfinite_step": first_any_nonfinite_step,
+ "first_any_nonfinite_groups": first_any_nonfinite,
+ "first_nonfinite_by_group": first_nonfinite_by_group,
+ "first_training_failure_step": first_training_failure_step,
+ "first_training_failure_groups": first_training_failure_groups,
+ "first_nonfinite_metrics": first_nonfinite_metrics,
"trajectory_metrics_finite": all(math.isfinite(value)
for value in numeric),
"validation_evaluations": 0,
@@ -209,8 +235,11 @@ def main():
handle.write("\n")
print(json.dumps({
"out": args.out,
- "first_nonfinite_step": first_nonfinite_step,
- "first_nonfinite_groups": first_nonfinite,
+ "first_any_nonfinite_step": first_any_nonfinite_step,
+ "first_any_nonfinite_groups": first_any_nonfinite,
+ "first_training_failure_step": first_training_failure_step,
+ "first_training_failure_groups": first_training_failure_groups,
+ "first_nonfinite_metrics": first_nonfinite_metrics,
"steps_recorded": len(trajectory),
}, indent=2))