diff --git a/src/pyrecest/evaluation/diagnostic_summaries.py b/src/pyrecest/evaluation/diagnostic_summaries.py index d2a48bf97..3fcf7eb39 100644 --- a/src/pyrecest/evaluation/diagnostic_summaries.py +++ b/src/pyrecest/evaluation/diagnostic_summaries.py @@ -252,14 +252,24 @@ def worst_time_windows( ], dtype=float, ) + error_scale = float(np.max(np.abs(errors))) + if error_scale == 0.0: + rmse = 0.0 + mae = 0.0 + p95 = 0.0 + else: + scaled_errors = errors / error_scale + rmse = error_scale * float(np.sqrt(np.mean(scaled_errors**2))) + mae = error_scale * float(np.mean(np.abs(scaled_errors))) + p95 = error_scale * float(np.percentile(scaled_errors, 95)) rows.append( { "time_start_s": float(start), "time_end_s": float(start + window_s), "count": int(errors.size), - "rmse": float(np.sqrt(np.mean(errors**2))), - "mae": float(np.mean(np.abs(errors))), - "p95": float(np.percentile(errors, 95)), + "rmse": rmse, + "mae": mae, + "p95": p95, "max": float(np.max(errors)), "mean_residual": ( None if residuals.size == 0 else float(np.mean(residuals)) diff --git a/tests/evaluation/test_diagnostic_summaries_extreme_errors.py b/tests/evaluation/test_diagnostic_summaries_extreme_errors.py new file mode 100644 index 000000000..1122cadab --- /dev/null +++ b/tests/evaluation/test_diagnostic_summaries_extreme_errors.py @@ -0,0 +1,22 @@ +import numpy as np + +from pyrecest.evaluation.diagnostic_summaries import worst_time_windows + + +def test_worst_time_windows_stays_finite_for_large_finite_errors(): + records = [ + {"time_s": 0.0, "error": -1.0e308}, + {"time_s": 1.0, "error": 1.0e308}, + ] + + with np.errstate(over="raise", invalid="raise"): + rows = worst_time_windows(records, window_s=5.0, top_n=1) + + assert len(rows) == 1 + row = rows[0] + assert np.isfinite(row["rmse"]) + assert np.isfinite(row["mae"]) + assert np.isfinite(row["p95"]) + np.testing.assert_allclose(row["rmse"], 1.0e308, rtol=1.0e-15) + np.testing.assert_allclose(row["mae"], 1.0e308, rtol=1.0e-15) + np.testing.assert_allclose(row["p95"], 9.0e307, rtol=1.0e-15)