Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 22 additions & 7 deletions ignite/metrics/regression/pearson_correlation.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from collections.abc import Callable
from collections.abc import Callable, Mapping

import torch

Expand Down Expand Up @@ -78,22 +78,36 @@ 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()
# 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)
self._sum_of_ys += y.sum().to(self._device)
self._sum_of_y_pred_squares += y_pred.square().sum().to(self._device)
self._sum_of_y_squares += y.square().sum().to(self._device)
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",
Expand Down Expand Up @@ -121,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())

55 changes: 55 additions & 0 deletions tests/ignite/metrics/regression/test_pearson_correlation.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,60 @@ 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_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)
Expand Down Expand Up @@ -263,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}"

Loading