From 705685c9122eeecc53e57285e44598c3453acb60 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Tue, 25 Nov 2025 09:40:07 +0100 Subject: [PATCH 1/2] fix v1 var calculation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. Masked variance calculation (lines 95-107): Changed from the numerically unstable E[X²] - E[X]² formula to the stable centered formula E[(X-μ)²] 2. Sigma clamping (line 609): Changed from torch.where(sigma < tolerance, 1.0, sigma) to torch.clamp(sigma, min=tolerance) --- v1/src/timesfm/pytorch_patched_decoder.py | 30 ++++++----------------- 1 file changed, 8 insertions(+), 22 deletions(-) diff --git a/v1/src/timesfm/pytorch_patched_decoder.py b/v1/src/timesfm/pytorch_patched_decoder.py index 15bf428..eb98bd5 100644 --- a/v1/src/timesfm/pytorch_patched_decoder.py +++ b/v1/src/timesfm/pytorch_patched_decoder.py @@ -94,26 +94,16 @@ def _masked_mean_std( # Calculate the number of valid elements num_valid_elements = torch.sum(mask, dim=1) - num_valid_elements = torch.where( - num_valid_elements == 0, - torch.tensor(1, - dtype=num_valid_elements.dtype, - device=num_valid_elements.device), - num_valid_elements, - ) + num_valid_elements = torch.clamp(num_valid_elements, min=1.0) - # Calculate the masked sum and squared sum + # Calculate the masked sum and mean masked_sum = torch.sum(arr * mask, dim=1) - masked_squared_sum = torch.sum((arr * mask)**2, dim=1) - - # Calculate the masked mean and standard deviation masked_mean = masked_sum / num_valid_elements - masked_var = masked_squared_sum / num_valid_elements - masked_mean**2 - masked_var = torch.where( - masked_var < 0.0, - torch.tensor(0.0, dtype=masked_var.dtype, device=masked_var.device), - masked_var, - ) + + # Calculate the masked variance using centered values (numerically stable) + masked_centered_arr = (arr - masked_mean.unsqueeze(-1)) * mask + masked_var = torch.sum(masked_centered_arr**2, dim=1) / num_valid_elements + masked_var = torch.clamp(masked_var, min=0.0) masked_std = torch.sqrt(masked_var) return masked_mean, masked_std @@ -616,11 +606,7 @@ class PatchedTimeSeriesDecoder(nn.Module): ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: """Input is of shape [B, N, P].""" mu, sigma = _masked_mean_std(inputs, patched_pads) - sigma = torch.where( - sigma < self.config.tolerance, - torch.tensor(1.0, dtype=sigma.dtype, device=sigma.device), - sigma, - ) + sigma = torch.clamp(sigma, min=self.config.tolerance) # Normalize each patch outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] From 177ab03e9c00bb8de5c75651029b8983ce44d799 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Wed, 18 Feb 2026 14:50:56 +0100 Subject: [PATCH 2/2] fix v1 jax version --- v1/src/timesfm/patched_decoder.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/v1/src/timesfm/patched_decoder.py b/v1/src/timesfm/patched_decoder.py index 48f8afe..fef9661 100644 --- a/v1/src/timesfm/patched_decoder.py +++ b/v1/src/timesfm/patched_decoder.py @@ -190,13 +190,13 @@ def _masked_mean_std(inputs: JTensor, num_valid_elements = jnp.where(num_valid_elements == 0, 1, num_valid_elements) - # Calculate the masked sum and squared sum of M + # Calculate the masked sum for mean and centered squared sum for variance. masked_sum = jnp.sum(arr * mask, axis=1) - masked_squared_sum = jnp.sum((arr * mask)**2, axis=1) # Calculate the masked mean and standard deviation masked_mean = masked_sum / num_valid_elements - masked_var = masked_squared_sum / num_valid_elements - masked_mean**2 + centered = (arr - masked_mean[:, None]) * mask + masked_var = jnp.sum(centered**2, axis=1) / num_valid_elements masked_var = jnp.where(masked_var < 0.0, 0.0, masked_var) masked_std = jnp.sqrt(masked_var) @@ -295,7 +295,7 @@ class PatchedTimeSeriesDecoder(base_layer.BaseLayer): patched_pads: JTensor) -> Tuple[JTensor, Tuple[JTensor, JTensor]]: """Input is of shape [B, N, P].""" mu, sigma = _masked_mean_std(inputs, patched_pads) - sigma = jnp.where(sigma < _TOLERANCE, 1.0, sigma) + sigma = jnp.maximum(sigma, _TOLERANCE) # Normalize each patch. outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] outputs = jnp.where(