diff --git a/src/gluonts/torch/util.py b/src/gluonts/torch/util.py index f4163dc215..76342ac4e3 100644 --- a/src/gluonts/torch/util.py +++ b/src/gluonts/torch/util.py @@ -95,10 +95,13 @@ def weighted_average( weights != 0, x * weights, torch.zeros_like(x) ) sum_weights = torch.clamp( - weights.sum(dim=dim) if dim else weights.sum(), min=1.0 + weights.sum(dim=dim) if dim is not None else weights.sum(), + min=1.0, ) return ( - weighted_tensor.sum(dim=dim) if dim else weighted_tensor.sum() + weighted_tensor.sum(dim=dim) + if dim is not None + else weighted_tensor.sum() ) / sum_weights else: return x.mean(dim=dim) diff --git a/test/torch/test_torch_util.py b/test/torch/test_torch_util.py index 90c6475bd7..b77a665b83 100644 --- a/test/torch/test_torch_util.py +++ b/test/torch/test_torch_util.py @@ -19,9 +19,19 @@ from gluonts.torch.util import ( lagged_sequence_values, unsqueeze_expand, + weighted_average, ) +def test_weighted_average_dim_zero(): + x = torch.tensor([[1.0, 10.0], [3.0, 30.0]]) + weights = torch.tensor([[1.0, 1.0], [3.0, 1.0]]) + + result = weighted_average(x, weights=weights, dim=0) + + torch.testing.assert_close(result, torch.tensor([2.5, 20.0])) + + @pytest.mark.parametrize( "lag_indices, prior_sequence, sequence, output_shape", [