fix v1 var calculation

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)
This commit is contained in:
Kashif Rasul
2025-11-25 09:40:07 +01:00
parent bf88c5dc88
commit 705685c912
+8 -22
View File
@@ -94,26 +94,16 @@ def _masked_mean_std(
# Calculate the number of valid elements # Calculate the number of valid elements
num_valid_elements = torch.sum(mask, dim=1) num_valid_elements = torch.sum(mask, dim=1)
num_valid_elements = torch.where( num_valid_elements = torch.clamp(num_valid_elements, min=1.0)
num_valid_elements == 0,
torch.tensor(1,
dtype=num_valid_elements.dtype,
device=num_valid_elements.device),
num_valid_elements,
)
# Calculate the masked sum and squared sum # Calculate the masked sum and mean
masked_sum = torch.sum(arr * mask, dim=1) 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_mean = masked_sum / num_valid_elements
masked_var = masked_squared_sum / num_valid_elements - masked_mean**2
masked_var = torch.where( # Calculate the masked variance using centered values (numerically stable)
masked_var < 0.0, masked_centered_arr = (arr - masked_mean.unsqueeze(-1)) * mask
torch.tensor(0.0, dtype=masked_var.dtype, device=masked_var.device), masked_var = torch.sum(masked_centered_arr**2, dim=1) / num_valid_elements
masked_var, masked_var = torch.clamp(masked_var, min=0.0)
)
masked_std = torch.sqrt(masked_var) masked_std = torch.sqrt(masked_var)
return masked_mean, masked_std return masked_mean, masked_std
@@ -616,11 +606,7 @@ class PatchedTimeSeriesDecoder(nn.Module):
) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
"""Input is of shape [B, N, P].""" """Input is of shape [B, N, P]."""
mu, sigma = _masked_mean_std(inputs, patched_pads) mu, sigma = _masked_mean_std(inputs, patched_pads)
sigma = torch.where( sigma = torch.clamp(sigma, min=self.config.tolerance)
sigma < self.config.tolerance,
torch.tensor(1.0, dtype=sigma.dtype, device=sigma.device),
sigma,
)
# Normalize each patch # Normalize each patch
outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] outputs = (inputs - mu[:, None, None]) / sigma[:, None, None]