From a132dade776e4c7b699e0385a354fedb3fad4e70 Mon Sep 17 00:00:00 2001 From: Tejas Date: Wed, 15 Apr 2026 17:11:13 -0400 Subject: [PATCH 1/3] fix: use float64 accumulators in PearsonCorrelation to prevent catastrophic cancellation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The naive E[X²] - (E[X])² formula loses all precision when values have large magnitude relative to their variance: both terms are ~μ² ≈ 1e16 while their difference (the variance) is ~O(1), which falls below float32's unit in the last place at that scale. Switch all five accumulators and the incoming batches to float64 on non-MPS devices. MPS does not support float64 and keeps the previous float32 behaviour. The final result is still returned as a Python float, so the public API is unchanged. Fixes #3662 Co-Authored-By: Claude Sonnet 4.6 --- .../metrics/regression/pearson_correlation.py | 17 +++++++++++------ .../regression/test_pearson_correlation.py | 18 ++++++++++++++++++ 2 files changed, 29 insertions(+), 6 deletions(-) diff --git a/ignite/metrics/regression/pearson_correlation.py b/ignite/metrics/regression/pearson_correlation.py index 1b4e8fa5d04b..8002f4d25c6d 100644 --- a/ignite/metrics/regression/pearson_correlation.py +++ b/ignite/metrics/regression/pearson_correlation.py @@ -78,15 +78,20 @@ def __init__( @reinit__is_reduced def reset(self) -> None: - self._sum_of_y_preds = torch.tensor(0.0, device=self._device) - self._sum_of_ys = torch.tensor(0.0, device=self._device) - self._sum_of_y_pred_squares = torch.tensor(0.0, device=self._device) - self._sum_of_y_squares = torch.tensor(0.0, device=self._device) - self._sum_of_products = torch.tensor(0.0, device=self._device) + # Use float64 accumulators to avoid catastrophic cancellation in + # E[X^2] - (E[X])^2 when values have large magnitude. MPS does not + # support float64, so fall back to float32 there. + acc_dtype = torch.float64 if self._device.type != "mps" else torch.float32 + self._sum_of_y_preds = torch.tensor(0.0, dtype=acc_dtype, device=self._device) + self._sum_of_ys = torch.tensor(0.0, dtype=acc_dtype, device=self._device) + self._sum_of_y_pred_squares = torch.tensor(0.0, dtype=acc_dtype, device=self._device) + self._sum_of_y_squares = torch.tensor(0.0, dtype=acc_dtype, device=self._device) + self._sum_of_products = torch.tensor(0.0, dtype=acc_dtype, device=self._device) self._num_examples = 0 def _update(self, output: tuple[torch.Tensor, torch.Tensor]) -> None: - y_pred, y = output[0].detach(), output[1].detach() + y_pred = output[0].detach().to(dtype=self._sum_of_y_preds.dtype) + y = output[1].detach().to(dtype=self._sum_of_y_preds.dtype) self._sum_of_y_preds += y_pred.sum().to(self._device) self._sum_of_ys += y.sum().to(self._device) self._sum_of_y_pred_squares += y_pred.square().sum().to(self._device) diff --git a/tests/ignite/metrics/regression/test_pearson_correlation.py b/tests/ignite/metrics/regression/test_pearson_correlation.py index 599e6fae2033..d88eaa4cadb7 100644 --- a/tests/ignite/metrics/regression/test_pearson_correlation.py +++ b/tests/ignite/metrics/regression/test_pearson_correlation.py @@ -137,6 +137,24 @@ def update_fn(engine: Engine, batch): assert pytest.approx(np_ans, rel=2e-4) == corr +def test_numerical_stability_large_offset(): + # float32 accumulators suffer catastrophic cancellation in E[X^2]-(E[X])^2 + # when values have large magnitude relative to their variance: both E[X^2] + # and (E[X])^2 are ~1e16 but their difference (the variance) is ~1, which + # falls below float32's ULP at that scale. float64 accumulators preserve + # the precision. MPS is excluded because it does not support float64. + offset = 1e8 + # y = 2*y_pred => perfect positive correlation; expected r = 1.0 + y_pred = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float64) + offset + y = torch.tensor([2.0, 4.0, 6.0, 8.0, 10.0], dtype=torch.float64) + offset + + m = PearsonCorrelation() # CPU device (float64 accumulators) + m.update((y_pred, y)) + result = m.compute() + + assert pytest.approx(1.0, abs=1e-6) == result + + def test_accumulator_detached(available_device): corr = PearsonCorrelation(device=available_device) assert corr._device == torch.device(available_device) From cd4088e172ead24d6fb17ab3d8a64cd080eaa68c Mon Sep 17 00:00:00 2001 From: Tejas Attarde <201814263+tejas-ae@users.noreply.github.com> Date: Thu, 24 Sep 2026 19:21:59 -0400 Subject: [PATCH 2/3] fix: retain Pearson accumulator dtype when loading old state --- ignite/metrics/regression/pearson_correlation.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/ignite/metrics/regression/pearson_correlation.py b/ignite/metrics/regression/pearson_correlation.py index 8002f4d25c6d..31e596b3b8a3 100644 --- a/ignite/metrics/regression/pearson_correlation.py +++ b/ignite/metrics/regression/pearson_correlation.py @@ -1,4 +1,4 @@ -from collections.abc import Callable +from collections.abc import Callable, Mapping import torch @@ -90,6 +90,7 @@ def reset(self) -> None: self._num_examples = 0 def _update(self, output: tuple[torch.Tensor, torch.Tensor]) -> None: + # Cast before square/product/reduction; widening their float32 results is too late. y_pred = output[0].detach().to(dtype=self._sum_of_y_preds.dtype) y = output[1].detach().to(dtype=self._sum_of_y_preds.dtype) self._sum_of_y_preds += y_pred.sum().to(self._device) @@ -99,6 +100,14 @@ def _update(self, output: tuple[torch.Tensor, torch.Tensor]) -> None: self._sum_of_products += (y_pred * y).sum().to(self._device) self._num_examples += y.shape[0] + def _load_state_dict_per_rank(self, state_dict: Mapping) -> None: + # Older checkpoints contain float32 sums. Keep the saved values while + # restoring the dtype chosen by reset() for this device. + acc_dtype = self._sum_of_y_preds.dtype + super()._load_state_dict_per_rank(state_dict) + for key in self._state_dict_all_req_keys[:-1]: + setattr(self, key, getattr(self, key).to(dtype=acc_dtype)) + @sync_all_reduce( "_sum_of_y_preds", "_sum_of_ys", @@ -126,3 +135,4 @@ def compute(self) -> float: r = cov / torch.clamp(torch.sqrt(y_pred_var * y_var), min=self.eps) return float(r.item()) + From c4a3b5ea327c7ffb8eb128fafb4a2bebb444207f Mon Sep 17 00:00:00 2001 From: Tejas Attarde <201814263+tejas-ae@users.noreply.github.com> Date: Thu, 24 Sep 2026 19:22:05 -0400 Subject: [PATCH 3/3] test: cover float32 Pearson inputs and legacy state --- .../regression/test_pearson_correlation.py | 37 +++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/tests/ignite/metrics/regression/test_pearson_correlation.py b/tests/ignite/metrics/regression/test_pearson_correlation.py index d88eaa4cadb7..85fdf42ba59a 100644 --- a/tests/ignite/metrics/regression/test_pearson_correlation.py +++ b/tests/ignite/metrics/regression/test_pearson_correlation.py @@ -155,6 +155,42 @@ def test_numerical_stability_large_offset(): assert pytest.approx(1.0, abs=1e-6) == result +def test_numerical_stability_float32_inputs(): + y_pred = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) + 1e4 + y = torch.tensor([2.0, 4.0, 6.0, 8.0, 10.0], dtype=torch.float32) + 1e4 + + m = PearsonCorrelation() + m.update((y_pred, y)) + + assert m.compute() == pytest.approx(1.0, abs=1e-6) + + +@pytest.mark.parametrize("with_prior_update", [False, True]) +def test_load_legacy_float32_state(with_prior_update): + original = PearsonCorrelation() + if with_prior_update: + original.update((torch.tensor([1.0, 2.0, 3.0]), torch.tensor([2.0, 4.0, 6.0]))) + + state = original.state_dict() + per_rank_state = next(iter(state.values()))[0] + accumulator_keys = original._state_dict_all_req_keys[:-1] + for key in accumulator_keys: + per_rank_state[key] = per_rank_state[key].float() + + restored = PearsonCorrelation() + restored.load_state_dict(state) + for key in accumulator_keys: + assert getattr(restored, key).dtype == torch.float64 + assert getattr(restored, key) == getattr(original, key) + assert restored._num_examples == original._num_examples + + y_pred = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype=torch.float32) + 1e4 + y = torch.tensor([2.0, 4.0, 6.0, 8.0, 10.0], dtype=torch.float32) + 1e4 + original.update((y_pred, y)) + restored.update((y_pred, y)) + assert restored.compute() == pytest.approx(original.compute(), abs=1e-6) + + def test_accumulator_detached(available_device): corr = PearsonCorrelation(device=available_device) assert corr._device == torch.device(available_device) @@ -281,3 +317,4 @@ def test_accumulator_device(self): ) for dev in devices: assert dev == metric_device, f"{type(dev)}:{dev} vs {type(metric_device)}:{metric_device}" +