diff options
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 11 |
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, |
