summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:35:24 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:35:24 -0500
commite9962317b3d976fe25e28a5693f2757e3758480a (patch)
tree78434f9b8f4e8c97293388244bcf4933e08da220
parentd2f7f35f112e5e21515c134548d6d27ada5a0b4f (diff)
diagnostic: distinguish RMS overflow from active failure
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py11
1 files changed, 9 insertions, 2 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py
index c144062..a045a75 100644
--- a/experiments/diagnose_kp_traffic_nonfinite.py
+++ b/experiments/diagnose_kp_traffic_nonfinite.py
@@ -153,6 +153,7 @@ def main():
first_training_failure_step = None
first_training_failure_groups = {}
first_nonfinite_metrics = []
+ first_nonfinite_by_metric = {}
for step, (x, y) in enumerate(train):
if step >= args.max_steps:
break
@@ -174,6 +175,11 @@ def main():
if (result["predictor_mse"] is not None
and not math.isfinite(float(result["predictor_mse"]))):
metric_names.append("predictor_mse")
+ for name in metric_names:
+ first_nonfinite_by_metric.setdefault(name, step + 1)
+ active_metric_names = [
+ name for name in metric_names if name in (
+ "loss", "teaching_rms", "instruction_rms", "predictor_mse")]
active_bad = {name: indices for name, indices in bad.items()
if name not in TRAINING_INACTIVE_GROUPS}
trajectory.append({
@@ -188,10 +194,10 @@ def main():
"predictor_mse": result["predictor_mse"],
"parameter_state": state,
})
- if active_bad or metric_names:
+ if active_bad or active_metric_names:
first_training_failure_step = step + 1
first_training_failure_groups = active_bad
- first_nonfinite_metrics = metric_names
+ first_nonfinite_metrics = active_metric_names
break
numeric = []
@@ -221,6 +227,7 @@ def main():
"first_any_nonfinite_step": first_any_nonfinite_step,
"first_any_nonfinite_groups": first_any_nonfinite,
"first_nonfinite_by_group": first_nonfinite_by_group,
+ "first_nonfinite_by_metric": first_nonfinite_by_metric,
"first_training_failure_step": first_training_failure_step,
"first_training_failure_groups": first_training_failure_groups,
"first_nonfinite_metrics": first_nonfinite_metrics,