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
5 changes: 3 additions & 2 deletions src/timesfm/flax/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
5 changes: 3 additions & 2 deletions src/timesfm/torch/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
4 changes: 2 additions & 2 deletions src/timesfm3/mlx/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
16 changes: 16 additions & 0 deletions src/timesfm3/torch/primitives_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 3 additions & 2 deletions src/timesfm3/torch/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
27 changes: 27 additions & 0 deletions tests/test_torch_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
Loading