diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:33:44 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:33:44 -0500 |
| commit | d2f7f35f112e5e21515c134548d6d27ada5a0b4f (patch) | |
| tree | 79813f49ca7e4df5fb6261785204971b5d9f3fba /experiments/diagnose_kp_traffic_nonfinite.py | |
| parent | efe356b3d9b99aa397caf8e8d70320dedc8d4450 (diff) | |
diagnostic: separate training-active nonfinite state
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 47 |
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)) |
