From 505c1072aa0360ad27a56710ed2f4fe6f1439b47 Mon Sep 17 00:00:00 2001 From: Steve Chan Date: Thu, 23 Jul 2026 05:18:26 +0000 Subject: [PATCH] fix(aggregation): fix nan_stategy=disable dropping weights --- src/torchmetrics/aggregation.py | 2 +- tests/unittests/bases/test_aggregation.py | 8 ++++++++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/src/torchmetrics/aggregation.py b/src/torchmetrics/aggregation.py index 1744df6f66b..7d0c97ed0ca 100644 --- a/src/torchmetrics/aggregation.py +++ b/src/torchmetrics/aggregation.py @@ -103,7 +103,7 @@ def _cast_and_nan_check_input( raise ValueError(f"`nan_strategy` shall be float but you pass {self.nan_strategy}") x[nans | nans_weight] = self.nan_strategy weight[nans | nans_weight] = 1 - else: + elif weight is None: weight = torch.ones_like(x) return x.to(self.dtype), weight.to(self.dtype) diff --git a/tests/unittests/bases/test_aggregation.py b/tests/unittests/bases/test_aggregation.py index 182e57260a5..6ffc80d000c 100644 --- a/tests/unittests/bases/test_aggregation.py +++ b/tests/unittests/bases/test_aggregation.py @@ -163,6 +163,14 @@ def test_error_on_wrong_nan_strategy(metric_class): metric_class(nan_strategy=[]) +def test_disable_preserves_weight(): + """`nan_strategy="disable"` should only skip nan checks, not also discarding `weight`.""" + metric = MeanMetric(nan_strategy="disable") + metric.update(torch.tensor(1.0), weight=torch.tensor(3.0)) + metric.update(torch.tensor(0.0), weight=torch.tensor(1.0)) + assert torch.allclose(metric.compute(), torch.tensor(0.75)) + + @pytest.mark.skipif(not hasattr(torch, "broadcast_to"), reason="PyTorch <1.8 does not have broadcast_to") @pytest.mark.parametrize( ("weights", "expected"), [(1, 11.5), (torch.ones(2, 1, 1), 11.5), (torch.tensor([1, 2]).reshape(2, 1, 1), 13.5)]