From 37cce368e46dc6462560066e5608a884d13f884c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8C=BA=E6=A2=93=E7=81=8F?= <116372750+Nicholas022400701@users.noreply.github.com> Date: Fri, 18 Sep 2026 21:29:15 +0800 Subject: [PATCH 1/3] Fix Variance ignoring dim in ensemble_metrics MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Variance.__call__ took the sample count from the leading dimension and subtracted a mean that only broadcast along the leading dimension, so a dim other than 0 either raised a broadcast error or returned a wrong variance without warning. _update_var had the same broadcast problem for a batch_dim other than 0. Signed-off-by: 区梓灏 <116372750+Nicholas022400701@users.noreply.github.com> --- physicsnemo/metrics/general/ensemble_metrics.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/physicsnemo/metrics/general/ensemble_metrics.py b/physicsnemo/metrics/general/ensemble_metrics.py index deaea01af8..350f841ebb 100644 --- a/physicsnemo/metrics/general/ensemble_metrics.py +++ b/physicsnemo/metrics/general/ensemble_metrics.py @@ -270,7 +270,10 @@ def _update_var( temp_n = inputs.shape[batch_dim] temp_sum = torch.sum(inputs, dim=batch_dim) - temp_sum2 = torch.sum((inputs - temp_sum / temp_n) ** 2, dim=batch_dim) + # Put the reduced dimension back so the mean broadcasts against inputs for + # any batch_dim, not only for the leading one. + temp_mean = torch.unsqueeze(temp_sum, batch_dim) / temp_n + temp_sum2 = torch.sum((inputs - temp_mean) ** 2, dim=batch_dim) delta = old_sum * temp_n / old_n - temp_sum @@ -327,7 +330,7 @@ def __call__(self, inputs: Tensor, dim: int = 0) -> Tensor: f"Input device, {inputs.device}, and Module device, {self.device}, must be the same." ) self.sum = torch.sum(inputs, dim=dim) - self.n = torch.as_tensor([inputs.shape[0]], device=self.device) + self.n = torch.as_tensor([inputs.shape[dim]], device=self.device) if ( DistributedManager.is_initialized() and dist.is_initialized() @@ -336,10 +339,14 @@ def __call__(self, inputs: Tensor, dim: int = 0) -> Tensor: dist.all_reduce(self.sum, op=dist.ReduceOp.SUM) dist.all_reduce(self.n, op=dist.ReduceOp.SUM) - self.sum2 = torch.sum((inputs - self.sum / self.n) ** 2, dim=dim) + # Put the reduced dimension back so the mean broadcasts against + # inputs for any dim, not only for the leading one. + mean = torch.unsqueeze(self.sum, dim) / self.n + self.sum2 = torch.sum((inputs - mean) ** 2, dim=dim) dist.all_reduce(self.sum2, op=dist.ReduceOp.SUM) else: - self.sum2 = torch.sum((inputs - self.sum / self.n) ** 2, dim=dim) + mean = torch.unsqueeze(self.sum, dim) / self.n + self.sum2 = torch.sum((inputs - mean) ** 2, dim=dim) if self.n < 2.0: return self.sum2 From ab033a9620a0df27a96b804618b3a51aae5ace36 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8C=BA=E6=A2=93=E7=81=8F?= <116372750+Nicholas022400701@users.noreply.github.com> Date: Fri, 18 Sep 2026 21:29:43 +0800 Subject: [PATCH 2/3] Test Variance and _update_var on a non leading dimension MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The existing test_means_var only covers dim 0 and is skipped without CUDA, so the broken dim handling never ran in CI. This test runs on CPU and covers a shape that used to raise a broadcast error and a square shape that used to return a wrong variance without an error. Signed-off-by: 区梓灏 <116372750+Nicholas022400701@users.noreply.github.com> --- test/metrics/test_ensemble_metrics_dim.py | 51 +++++++++++++++++++++++ 1 file changed, 51 insertions(+) create mode 100644 test/metrics/test_ensemble_metrics_dim.py diff --git a/test/metrics/test_ensemble_metrics_dim.py b/test/metrics/test_ensemble_metrics_dim.py new file mode 100644 index 0000000000..211fbba478 --- /dev/null +++ b/test/metrics/test_ensemble_metrics_dim.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023 - 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import torch + +import physicsnemo.metrics.general.ensemble_metrics as em + + +@pytest.mark.parametrize("dim", [1, 2, -1]) +@pytest.mark.parametrize("input_shape", [(4, 6, 5), (6, 6, 5)]) +def test_variance_dim( + device, input_shape, dim, rtol: float = 1e-4, atol: float = 1e-4 +): + # The ensemble dimension is not the leading one here. Variance has to take + # the sample count from that dimension and subtract a mean that broadcasts + # along it. The square shape is the case that used to run without an error + # and return a wrong value, the other shape used to raise a broadcast error. + x = torch.randn(input_shape, device=device) + remaining_shape = list(input_shape) + del remaining_shape[dim] + expected_n = input_shape[dim] + + V = em.Variance(remaining_shape, device=device) + var = V(x, dim=dim) + assert var.shape == tuple(remaining_shape) + assert V.n.item() == expected_n + assert torch.allclose(var, torch.var(x, dim=dim), rtol=rtol, atol=atol) + assert torch.allclose(V.mean, torch.mean(x, dim=dim), rtol=rtol, atol=atol) + + # The standalone update rule has to accept the same batch dimension. + _sum, _sum2, _n = em._update_var(V.sum, V.sum2, V.n, x, batch_dim=dim) + doubled = torch.cat((x, x), dim=dim) + assert _n.item() == 2 * expected_n + assert torch.allclose(_sum / _n, doubled.mean(dim=dim), rtol=rtol, atol=atol) + assert torch.allclose( + _sum2 / (_n - 1.0), doubled.var(dim=dim), rtol=rtol, atol=atol + ) From c8e89d4531c2c7e232f11b3cf4a9d4e95e807361 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8C=BA=E6=A2=93=E7=81=8F?= <116372750+Nicholas022400701@users.noreply.github.com> Date: Fri, 18 Sep 2026 21:37:19 +0800 Subject: [PATCH 3/3] Apply ruff format to the new test MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: 区梓灏 <116372750+Nicholas022400701@users.noreply.github.com> --- test/metrics/test_ensemble_metrics_dim.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/test/metrics/test_ensemble_metrics_dim.py b/test/metrics/test_ensemble_metrics_dim.py index 211fbba478..5f6f88b639 100644 --- a/test/metrics/test_ensemble_metrics_dim.py +++ b/test/metrics/test_ensemble_metrics_dim.py @@ -22,9 +22,7 @@ @pytest.mark.parametrize("dim", [1, 2, -1]) @pytest.mark.parametrize("input_shape", [(4, 6, 5), (6, 6, 5)]) -def test_variance_dim( - device, input_shape, dim, rtol: float = 1e-4, atol: float = 1e-4 -): +def test_variance_dim(device, input_shape, dim, rtol: float = 1e-4, atol: float = 1e-4): # The ensemble dimension is not the leading one here. Variance has to take # the sample count from that dimension and subtract a mean that broadcasts # along it. The square shape is the case that used to run without an error