diff --git a/src/timesfm/flax/util.py b/src/timesfm/flax/util.py index ec70d724..ef792e99 100644 --- a/src/timesfm/flax/util.py +++ b/src/timesfm/flax/util.py @@ -101,7 +101,8 @@ def revin( elif len(mu.shape) == len(x.shape) - 2: mu = mu[..., None, None] sigma = sigma[..., None, None] + safe_sigma = jnp.where(sigma < _TOLERANCE, 1.0, sigma) if reverse: - return x * sigma + mu + return x * safe_sigma + mu else: - return (x - mu) / jnp.where(sigma < _TOLERANCE, 1.0, sigma) + return (x - mu) / safe_sigma diff --git a/src/timesfm/torch/util.py b/src/timesfm/torch/util.py index 81efe292..2bdf8eb4 100644 --- a/src/timesfm/torch/util.py +++ b/src/timesfm/torch/util.py @@ -88,7 +88,8 @@ def revin( mu = mu[..., None, None] sigma = sigma[..., None, None] + safe_sigma = torch.where(sigma < _TOLERANCE, 1.0, sigma) if reverse: - return x * sigma + mu + return x * safe_sigma + mu else: - return (x - mu) / torch.where(sigma < _TOLERANCE, 1.0, sigma) + return (x - mu) / safe_sigma diff --git a/src/timesfm3/mlx/util.py b/src/timesfm3/mlx/util.py index 1b87ea96..f0e36771 100644 --- a/src/timesfm3/mlx/util.py +++ b/src/timesfm3/mlx/util.py @@ -47,9 +47,9 @@ def revin( mu, sigma = mu[..., None], sigma[..., None] elif mu.ndim == x.ndim - 2: mu, sigma = mu[..., None, None], sigma[..., None, None] - if reverse: - return x * sigma + mu safe_sigma = mx.where(sigma < DIV_TOL, 1.0, sigma) + if reverse: + return x * safe_sigma + mu return (x - mu) / safe_sigma diff --git a/src/timesfm3/torch/primitives_test.py b/src/timesfm3/torch/primitives_test.py index 7d3a2f5b..cc4e471a 100644 --- a/src/timesfm3/torch/primitives_test.py +++ b/src/timesfm3/torch/primitives_test.py @@ -287,6 +287,22 @@ def test_revin_near_zero_sigma(self): expected = torch.tensor([[[0.0, 0.0, 0.0]]]) np.testing.assert_allclose(normalized_x.numpy(), expected.numpy(), atol=1e-6) + def test_revin_zero_sigma_roundtrip(self): + x = torch.tensor([[[4.5, 5.0, 5.5]]]) + mu = torch.tensor([[5.0]]) + sigma = torch.tensor([[0.0]]) + normalized_x = torch_util.revin(x, mu, sigma, reverse=False) + recovered_x = torch_util.revin(normalized_x, mu, sigma, reverse=True) + np.testing.assert_allclose(recovered_x.numpy(), x.numpy(), atol=1e-5) + + def test_revin_near_zero_sigma_roundtrip(self): + x = torch.tensor([[[1.0, 2.0, 3.0]]]) + mu = torch.tensor([[2.0]]) + sigma = torch.tensor([[1e-7]]) + normalized_x = torch_util.revin(x, mu, sigma, reverse=False) + recovered_x = torch_util.revin(normalized_x, mu, sigma, reverse=True) + np.testing.assert_allclose(recovered_x.numpy(), x.numpy(), atol=1e-5) + def test_get_output_patch_via_roll(self): x = torch.tensor([[[[1, 2], [3, 4], [5, 6], [7, 8]]]], dtype=torch.float32) rolls = 2 diff --git a/src/timesfm3/torch/util.py b/src/timesfm3/torch/util.py index 06350ba6..f482e216 100644 --- a/src/timesfm3/torch/util.py +++ b/src/timesfm3/torch/util.py @@ -251,10 +251,11 @@ def revin( sigma = sigma.unsqueeze(-1).unsqueeze(-1) else: raise ValueError(f"Unsupported shapes for x and mu: {x.shape}, {mu.shape}.") + safe_sigma = _make_safe_for_division(sigma) if reverse: - return x * sigma + mu + return x * safe_sigma + mu else: - return (x - mu) / _make_safe_for_division(sigma) + return (x - mu) / safe_sigma def get_output_patch_via_roll( diff --git a/tests/test_torch_utils.py b/tests/test_torch_utils.py index 76f2296a..bb674e72 100644 --- a/tests/test_torch_utils.py +++ b/tests/test_torch_utils.py @@ -254,6 +254,33 @@ def test_near_zero_sigma_guarded_by_tolerance(self): expected = torch.tensor([[-1.0, 0.0, 1.0]]) torch.testing.assert_close(normed, expected, atol=1e-5, rtol=1e-5) + def test_zero_sigma_roundtrip_identity(self): + """Zero variance must reconstruct the original input in roundtrip. + + When sigma is zero, forward normalization guards division by using + effective scale 1.0. Reverse denormalization must symmetrically use the + same effective scale 1.0 to preserve forecast deltas. + """ + x = torch.tensor([[4.5, 5.0, 5.5]]) + mu = torch.tensor([5.0]) + sigma = torch.tensor([0.0]) + + normed = revin(x, mu, sigma, reverse=False) + recovered = revin(normed, mu, sigma, reverse=True) + + torch.testing.assert_close(recovered, x, atol=1e-5, rtol=1e-5) + + def test_near_zero_sigma_roundtrip_identity(self): + """Near-zero variance below tolerance must reconstruct original input in roundtrip.""" + x = torch.tensor([[1.0, 2.0, 3.0]]) + mu = torch.tensor([2.0]) + sigma = torch.tensor([_TOLERANCE / 2]) + + normed = revin(x, mu, sigma, reverse=False) + recovered = revin(normed, mu, sigma, reverse=True) + + torch.testing.assert_close(recovered, x, atol=1e-5, rtol=1e-5) + def test_roundtrip_with_batched_3d_input(self): """RevIN must broadcast correctly for (batch, patches, patch_len) tensors — the actual shape used during patched decoding."""