diff --git a/sdm/models/timesfm3/configs.py b/sdm/models/timesfm3/configs.py new file mode 100644 index 000000000..8191e3371 --- /dev/null +++ b/sdm/models/timesfm3/configs.py @@ -0,0 +1,87 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import dataclass +from typing import Literal + + +@dataclass(frozen=True) +class ResidualBlockConfig: + """Configure a TimesFM-3 residual block. + + Args: + hidden_dims: Hidden-layer width. + output_dims: Output width. + use_bias: Whether linear layers use bias parameters. + activation: Hidden-layer activation. + identity_skip: Whether to use an identity residual connection. + prenorm: Normalization applied before the hidden layer. + """ + + hidden_dims: int + output_dims: int + use_bias: bool + activation: Literal["relu", "swish", "none"] + identity_skip: bool = False + prenorm: Literal["rms", "none"] = "none" + + +@dataclass(frozen=True) +class TransformerConfig: + """Configure a TimesFM-3 mixing transformer. + + Args: + model_dims: Input and output width. + hidden_dims: Feed-forward hidden width. + num_heads: Number of attention heads. + qk_norm: Query and key normalization. + use_bias: Whether linear layers use bias parameters. + use_rope_seq: Whether temporal attention uses rotary embeddings. + use_rope_var: Whether variate attention uses rotary embeddings. + ff_activation: Feed-forward activation. + v_norm: Value normalization. + causal_attention: Whether temporal attention is causal. + use_memory_efficient_attention: Whether to retain TimesFM's + square-root head-dimension logit scaling. + use_sdpa: Whether to use PyTorch scaled dot-product attention. + """ + + model_dims: int + hidden_dims: int + num_heads: int + qk_norm: Literal["rms", "none"] + use_bias: bool + use_rope_seq: bool + use_rope_var: bool + ff_activation: Literal["relu", "swish", "none"] + v_norm: Literal["rms", "none"] = "none" + causal_attention: bool = True + use_memory_efficient_attention: bool = True + use_sdpa: bool = True + + +@dataclass(frozen=True) +class StackedTransformersConfig: + """Configure a stack of TimesFM-3 mixing transformers. + + Args: + num_layers: Number of transformer layers. + transformer: Configuration shared by every layer. + """ + + num_layers: int + transformer: TransformerConfig diff --git a/sdm/models/timesfm3/core.py b/sdm/models/timesfm3/core.py new file mode 100644 index 000000000..eb5afac91 --- /dev/null +++ b/sdm/models/timesfm3/core.py @@ -0,0 +1,360 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any, cast + +import torch +from torch import Tensor + +from sdm.models.timesfm3.configs import ( + ResidualBlockConfig, + StackedTransformersConfig, + TransformerConfig, +) +from sdm.models.timesfm3.cpm_revin_refine import ( + cpm_iterative_revin_refine, +) +from sdm.models.timesfm3.dense import ResidualBlock +from sdm.models.timesfm3.transformer import StackedMixingTransformer +from sdm.models.timesfm3.util import ( + get_output_patch_via_roll, + get_running_stats, + revin, +) + + +class _TimesFM3Model(torch.nn.Module): + """Implement the Google TimesFM-3 full-sequence model. + + Args: + input_patch_len: Number of time steps in each input patch. + output_patch_len: Number of time steps predicted from each patch. + quantiles: Quantiles predicted by the output head. + residual_block_config: Pre-transformer residual-block configuration. + transformer_config: Transformer-stack configuration. + use_variate_attention: Whether to attend across variates. + value_clip: Absolute bound applied to inputs and predictions. + use_stitching: Whether decoding will stitch overlapping predictions. + use_linear_detrending: Whether decoding will remove linear trends. + linear_detrending_threshold: Ratio controlling linear detrending. + use_iterative_cpm_revin: Whether to refine RevIN statistics for + Contiguous Patch Masking positions. + use_frozen_running_stats: Whether decoding will freeze statistics at + the context boundary. + device: Device on which to create parameters and buffers. + dtype: Data type of parameters. + """ + + def __init__( + self, + input_patch_len: int = 32, + output_patch_len: int = 64, + quantiles: Sequence[float] | None = None, + residual_block_config: ResidualBlockConfig + | dict[str, Any] + | None = None, + transformer_config: StackedTransformersConfig + | dict[str, Any] + | None = None, + use_variate_attention: bool = True, + value_clip: float = 1e20, + use_stitching: bool = True, + use_linear_detrending: bool = True, + linear_detrending_threshold: float = 0.5, + use_iterative_cpm_revin: bool = True, + use_frozen_running_stats: bool = False, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + if quantiles is None: + quantiles = [index / 10 for index in range(1, 10)] + if residual_block_config is None: + residual_block_config = ResidualBlockConfig( + hidden_dims=1280, + output_dims=1280, + use_bias=False, + activation="relu", + ) + elif isinstance(residual_block_config, dict): + residual_block_config = ResidualBlockConfig( + **cast(dict[str, Any], residual_block_config) + ) + + if transformer_config is None: + transformer_config = StackedTransformersConfig( + num_layers=20, + transformer=TransformerConfig( + model_dims=1280, + hidden_dims=1280, + num_heads=16, + qk_norm="rms", + use_rope_seq=True, + use_rope_var=False, + use_bias=False, + ff_activation="relu", + ), + ) + elif isinstance(transformer_config, dict): + transformer_data = transformer_config.get("transformer", {}) + if isinstance(transformer_data, dict): + transformer_data = TransformerConfig( + **cast(dict[str, Any], transformer_data) + ) + stack_data = dict(transformer_config) + stack_data["transformer"] = transformer_data + transformer_config = StackedTransformersConfig(**stack_data) + + if output_patch_len % input_patch_len != 0: + raise ValueError( + f"Output patch length {output_patch_len} must be a multiple " + f"of input patch length {input_patch_len}." + ) + if ( + residual_block_config.output_dims + != transformer_config.transformer.model_dims + ): + raise ValueError( + "Residual-block output dimensions must match transformer " + "model dimensions." + ) + if use_stitching and output_patch_len <= input_patch_len: + raise ValueError( + "Stitching requires output_patch_len > input_patch_len." + ) + + self.input_patch_len = input_patch_len + self.output_patch_len = output_patch_len + self.quantiles = list(quantiles) + self.num_quantiles = len(self.quantiles) + self.rolls = output_patch_len // input_patch_len + self.residual_block_config = residual_block_config + self.transformer_config = transformer_config + self.use_variate_attention = use_variate_attention + self.value_clip = value_clip + self.use_stitching = use_stitching + self.use_linear_detrending = use_linear_detrending + self.linear_detrending_threshold = linear_detrending_threshold + self.use_iterative_cpm_revin = use_iterative_cpm_revin + self.use_frozen_running_stats = use_frozen_running_stats + + if use_stitching: + self._stitching_extract_len = min( + 2 * input_patch_len, + output_patch_len, + ) + + self.pre_transformer_resblock = ResidualBlock( + config=residual_block_config, + input_dims=2 * (input_patch_len + output_patch_len), + device=device, + dtype=dtype, + ) + self.transformer_stack = StackedMixingTransformer( + config=transformer_config, + use_variate_attention=use_variate_attention, + device=device, + dtype=dtype, + ) + self.output_head = torch.nn.Linear( + transformer_config.transformer.model_dims, + output_patch_len * self.num_quantiles, + bias=True, + device=device, + dtype=dtype, + ) + + def _preprocess( + self, + values: Tensor, + masks: Tensor, + patch_is_target: Tensor, + freeze_after: int | None = None, + patch_cpm_mask: Tensor | None = None, + ) -> tuple[ + Tensor, + Tensor, + Tensor, + tuple[Tensor, Tensor], + Tensor, + ]: + running_n, running_mean, running_std = get_running_stats( + values, + masks, + ) + if freeze_after is not None: + num_patches = values.shape[2] + if 0 <= freeze_after < num_patches - 1: + running_mean[:, :, freeze_after + 1 :] = running_mean[ + :, :, freeze_after : freeze_after + 1 + ] + running_std[:, :, freeze_after + 1 :] = running_std[ + :, :, freeze_after : freeze_after + 1 + ] + + if patch_cpm_mask is not None: + cpm_mask = patch_cpm_mask[:, None, :, None] + masks = masks | (cpm_mask & patch_is_target.unsqueeze(-1)) + + normalized_values = revin( + values, + running_mean, + running_std, + ) + normalized_values = torch.where(masks, 0.0, normalized_values) + + future_values, wrap_mask = get_output_patch_via_roll( + values, + self.rolls, + ) + future_values = revin( + future_values, + running_mean, + running_std, + ) + future_masks, _ = get_output_patch_via_roll(masks, self.rolls) + future_masks = future_masks | patch_is_target.unsqueeze(-1) | wrap_mask + future_values = torch.where(future_masks, 0.0, future_values) + + values_with_future = torch.cat( + [normalized_values, future_values], + dim=-1, + ) + masks_with_future = torch.cat([masks, future_masks], dim=-1) + input_dtype = self.pre_transformer_resblock.hidden_layer.weight.dtype + residual_input = torch.cat( + [ + values_with_future.to(input_dtype), + masks_with_future.to(input_dtype), + ], + dim=-1, + ) + transformer_input = self.pre_transformer_resblock(residual_input) + patch_mask = masks_with_future.all(dim=-1) + + return ( + residual_input, + transformer_input, + patch_mask, + (running_mean, running_std), + running_n, + ) + + def forward( + self, + values: Tensor, + masks: Tensor, + patch_is_target: Tensor, + *, + freeze_after: int | None = None, + patch_cpm_mask: Tensor | None = None, + return_aux_outputs: bool = False, + ) -> dict[str, Any]: + """Predict every quantile for each input patch. + + Args: + values: Patched series with shape ``[B, V, N, P]``. + masks: Invalid-value mask with shape ``[B, V, N, P]``. + patch_is_target: Target-patch indicator with shape ``[B, V, N]``. + freeze_after: Optional final patch included in running statistics. + patch_cpm_mask: Contiguous Patch Masking indicator with shape + ``[B, N]``. + return_aux_outputs: Whether to include intermediate tensors. + + Returns: + Mapping containing ``logits`` with shape ``[B, V, N, O, Q]`` and + the RevIN statistics used for denormalization. + """ + values = values.nan_to_num(nan=0.0).clamp( + -self.value_clip, + self.value_clip, + ) + masks = masks.bool() + if values.shape[-1] != self.input_patch_len: + raise ValueError( + f"Input patch length {values.shape[-1]} does not match " + f"configured length {self.input_patch_len}." + ) + + ( + residual_input, + transformer_input, + transformer_patch_mask, + revin_stats, + running_n, + ) = self._preprocess( + values, + masks, + patch_is_target, + freeze_after=freeze_after, + patch_cpm_mask=patch_cpm_mask, + ) + + effective_patch_mask = transformer_patch_mask.cummin(dim=2).values + transformer_output, attention_masks = self.transformer_stack( + transformer_input, + effective_patch_mask, + ) + raw_logits = self.output_head(transformer_output) + revin_mean, revin_std = revin_stats + + if self.use_iterative_cpm_revin and patch_cpm_mask is not None: + refined_mean, refined_std = cpm_iterative_revin_refine( + raw_logits, + revin_n=running_n, + revin_mu=revin_mean, + revin_sigma=revin_std, + patch_cpm_mask=patch_cpm_mask, + median_q_idx=self.num_quantiles // 2, + rolls=self.rolls, + patch_len=self.input_patch_len, + num_quantiles=self.num_quantiles, + value_clip=self.value_clip, + ) + cpm_mask = patch_cpm_mask.unsqueeze(1) + revin_mean = torch.where(cpm_mask, refined_mean, revin_mean) + revin_std = torch.where(cpm_mask, refined_std, revin_std) + + logits = revin( + raw_logits, + revin_mean, + revin_std, + reverse=True, + ).clamp(-self.value_clip, self.value_clip) + batch_size, num_variates, num_patches = logits.shape[:3] + logits = logits.view( + batch_size, + num_variates, + num_patches, + self.output_patch_len, + self.num_quantiles, + ) + + outputs: dict[str, Any] = { + "logits": logits, + "revin_stats": revin_stats, + } + if return_aux_outputs: + outputs["__call__:resblock_input"] = residual_input + outputs["__call__:transformer_input"] = transformer_input + outputs["__call__:seq_attn_mask"] = attention_masks + outputs["__call__:transformer_output"] = transformer_output + return outputs diff --git a/sdm/models/timesfm3/cpm_revin_refine.py b/sdm/models/timesfm3/cpm_revin_refine.py new file mode 100644 index 000000000..d5de3ed50 --- /dev/null +++ b/sdm/models/timesfm3/cpm_revin_refine.py @@ -0,0 +1,159 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import torch +from torch import Tensor + +from sdm.models.timesfm3.util import revin, update_running_stats + + +def cpm_iterative_revin_refine( + raw_logits: Tensor, + revin_n: Tensor, + revin_mu: Tensor, + revin_sigma: Tensor, + patch_cpm_mask: Tensor, + median_q_idx: int, + rolls: int, + patch_len: int, + num_quantiles: int, + value_clip: float = 1e9, +) -> tuple[Tensor, Tensor]: + """Refine normalization statistics across hidden time-series patches. + + A patch is a short span of time steps. RevIN (reversible instance + normalization) scales values using their mean and standard deviation, then + restores predictions to the original scale. Contiguous patch masking (CPM) + hides adjacent patches. For hidden patches, this function uses earlier + median predictions in place of missing values to update the statistics + (or zeros before any prediction is available). Updates carry across + output-patch blocks; visible patches keep their supplied statistics and + reset the run. + + Args: + raw_logits: Normalized predictions with shape ``[B, V, N, O * Q]``, + where ``B`` is the batch size, ``V`` is the number of variates, + ``N`` is the number of patches, ``O`` is the output patch length, + and ``Q`` is the number of quantiles. + revin_n: Running counts with shape ``[B, V, N]``. + revin_mu: Running means with shape ``[B, V, N]``. + revin_sigma: Running standard deviations with shape ``[B, V, N]``. + patch_cpm_mask: Boolean mask with shape ``[B, N]``; ``True`` marks + a hidden patch. + median_q_idx: Index of the median quantile used as the point estimate. + rolls: Number of input patches covered by each output patch. + patch_len: Input patch length. + num_quantiles: Number of predicted quantiles. + value_clip: Absolute bound applied to estimated values. + + Returns: + Refined means and standard deviations, each with shape ``[B, V, N]``. + """ + batch_size, num_variates, num_patches, _ = raw_logits.shape + device = raw_logits.device + + # [B, V, N, O * Q] -> [B, V, N, R, P] + median_logits = raw_logits.reshape( + batch_size, + num_variates, + num_patches, + rolls, + patch_len, + num_quantiles, + )[..., median_q_idx] + + carry_n = torch.zeros( + (batch_size, num_variates), + dtype=torch.float32, + device=device, + ) + carry_mu = torch.zeros_like(carry_n) + carry_sigma = torch.zeros_like(carry_n) + anchor_predicted_values = torch.zeros( + (batch_size, num_variates, rolls, patch_len), + dtype=torch.float32, + device=device, + ) + block_offset = torch.zeros( + batch_size, + dtype=torch.long, + device=device, + ) + step_masks = torch.zeros( + (batch_size, num_variates, patch_len), + dtype=torch.bool, + device=device, + ) + + refined_mu = [] + refined_sigma = [] + for index in range(num_patches): + actual_n = revin_n[:, :, index] + actual_mu = revin_mu[:, :, index] + actual_sigma = revin_sigma[:, :, index] + current_step_logits = median_logits[:, :, index] + is_cpm = patch_cpm_mask[:, index : index + 1] + + offset_index = block_offset.view(batch_size, 1, 1, 1).expand( + -1, + num_variates, + 1, + patch_len, + ) + predicted_values_step = anchor_predicted_values.gather( + 2, + offset_index, + ).squeeze(2) + new_n, new_mu, new_sigma = update_running_stats( + carry_n, + carry_mu, + carry_sigma, + predicted_values_step, + step_masks, + ) + + out_n = torch.where(is_cpm, new_n, actual_n) + out_mu = torch.where(is_cpm, new_mu, actual_mu) + out_sigma = torch.where(is_cpm, new_sigma, actual_sigma) + + new_block_offset = torch.where( + is_cpm.squeeze(-1), + (block_offset + 1) % rolls, + torch.zeros_like(block_offset), + ) + should_update_anchor = new_block_offset == 0 + + step_predicted_values = revin( + current_step_logits, + out_mu, + out_sigma, + reverse=True, + ).clamp(-value_clip, value_clip) + anchor_predicted_values = torch.where( + should_update_anchor.view(batch_size, 1, 1, 1), + step_predicted_values, + anchor_predicted_values, + ) + + carry_n = out_n + carry_mu = out_mu + carry_sigma = out_sigma + block_offset = new_block_offset + refined_mu.append(out_mu) + refined_sigma.append(out_sigma) + + return torch.stack(refined_mu, dim=2), torch.stack(refined_sigma, dim=2) diff --git a/sdm/models/timesfm3/dense.py b/sdm/models/timesfm3/dense.py new file mode 100644 index 000000000..62f4c3153 --- /dev/null +++ b/sdm/models/timesfm3/dense.py @@ -0,0 +1,104 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from typing import Any + +import torch +from torch import Tensor + +from sdm.models.timesfm3.configs import ResidualBlockConfig +from sdm.models.timesfm3.util import get_activation_fn + + +class ResidualBlock(torch.nn.Module): + """Apply the TimesFM-3 two-layer residual block. + + Args: + config: Residual block configuration. + input_dims: Size of the last input dimension. + device: Device on which to create parameters. + dtype: Data type of the parameters. + """ + + def __init__( + self, + config: ResidualBlockConfig, + input_dims: int, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + if config.identity_skip and config.output_dims != input_dims: + raise ValueError( + "identity_skip requires output_dims to match input_dims, got " + f"{config.output_dims} and {input_dims}" + ) + self.config = config + factory_kwargs: dict[str, Any] = {"device": device, "dtype": dtype} + + self.hidden_layer = torch.nn.Linear( + in_features=input_dims, + out_features=config.hidden_dims, + bias=config.use_bias, + **factory_kwargs, + ) + self.output_layer = torch.nn.Linear( + in_features=config.hidden_dims, + out_features=config.output_dims, + bias=config.use_bias, + **factory_kwargs, + ) + if config.identity_skip: + self.residual_layer: torch.nn.Linear | None = None + else: + self.residual_layer = torch.nn.Linear( + in_features=input_dims, + out_features=config.output_dims, + bias=config.use_bias, + **factory_kwargs, + ) + + self.activation = get_activation_fn(config.activation) + if config.prenorm == "rms": + self.pre_norm: torch.nn.RMSNorm | None = torch.nn.RMSNorm( + input_dims, + **factory_kwargs, + ) + elif config.prenorm == "none": + self.pre_norm = None + else: + raise AssertionError( + f"Unhandled pre-normalization: {config.prenorm}" + ) + + def forward(self, x: Tensor) -> Tensor: + """Transform input values and add the residual connection. + + Args: + x: Input values with shape ``[..., D]``, where ``D`` is the input + dimension. + + Returns: + Transformed values with shape ``[..., O]``, where ``O`` is the + configured output dimension. + """ + hidden_input = self.pre_norm(x) if self.pre_norm is not None else x + hidden_output = self.activation(self.hidden_layer(hidden_input)) + output = self.output_layer(hidden_output) + if self.residual_layer is not None: + return output + self.residual_layer(x) + return output + x diff --git a/sdm/models/timesfm3/normalization.py b/sdm/models/timesfm3/normalization.py new file mode 100644 index 000000000..cb95abe5e --- /dev/null +++ b/sdm/models/timesfm3/normalization.py @@ -0,0 +1,65 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import math +from typing import Any + +import torch +import torch.nn.functional as F +from torch import Tensor + +_RECIPROCAL_OF_SOFTPLUS_0 = 1.442695041 + + +class PerDimScale(torch.nn.Module): + r"""Apply learnable per-dimension scaling. + + Args: + num_dims: The number of dimensions. + device: The device. + dtype: The dtype. + """ + + def __init__( + self, + num_dims: int, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + factory_kwargs: dict[str, Any] = {"device": device, "dtype": dtype} + + self.num_dims = num_dims + self.per_dim_scale = torch.nn.Parameter( + torch.zeros(num_dims, **factory_kwargs) + ) + + def forward(self, tensor: Tensor) -> Tensor: + """Apply per-dimension scaling. + + Args: + tensor: Input with shape ``[..., D]``. + + Returns: + Scaled tensor with shape ``[..., D]``. + """ + return ( + tensor + * _RECIPROCAL_OF_SOFTPLUS_0 + / math.sqrt(self.num_dims) + * F.softplus(self.per_dim_scale) + ) diff --git a/sdm/models/timesfm3/transformer.py b/sdm/models/timesfm3/transformer.py new file mode 100644 index 000000000..3466bce0c --- /dev/null +++ b/sdm/models/timesfm3/transformer.py @@ -0,0 +1,520 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import math +from typing import Any, Literal + +import torch +import torch.nn.functional as F +from torch import Tensor + +from sdm.models.timesfm3.configs import ( + StackedTransformersConfig, + TransformerConfig, +) +from sdm.models.timesfm3.normalization import PerDimScale +from sdm.models.timesfm3.util import get_activation_fn + + +def make_attn_mask(patch_mask: Tensor, causal: bool = True) -> Tensor: + """Create an attention mask in which ``True`` permits attention. + + Args: + patch_mask: Masked patches with shape ``[B, N]``. + causal: Whether queries may attend only to preceding positions. + + Returns: + Boolean mask with shape ``[B, 1, N, N]`` when causal and broadcastable + shape ``[B, 1, 1, N]`` otherwise. + """ + mask = ~patch_mask[:, None, None, :] + if not causal: + return mask + causal_mask = torch.ones( + patch_mask.size(1), + patch_mask.size(1), + dtype=torch.bool, + device=patch_mask.device, + ).tril() + return causal_mask[None, None] & mask + + +class RotaryPositionalEmbedding(torch.nn.Module): + """Apply RoPE from the `RoFormer paper `_. + + Args: + embedding_dims: Size of the last input dimension. + device: Device on which to create the timescale buffer. + """ + + timescale: Tensor + + def __init__( + self, + embedding_dims: int, + device: torch.device | str | None = None, + ) -> None: + super().__init__() + self.embedding_dims = embedding_dims + self.register_buffer( + "timescale", + self._make_timescale(device), + persistent=False, + ) + + def _make_timescale(self, device: torch.device | str | None) -> Tensor: + half_dim = self.embedding_dims // 2 + fraction = ( + 2.0 + * torch.arange(half_dim, dtype=torch.float32, device=device) + / self.embedding_dims + ) + return 10_000**fraction + + def _reset_timescale(self, device: torch.device | str) -> None: + self.timescale = self._make_timescale(device) + + def forward(self, inputs: Tensor) -> Tensor: + """Apply rotary positional embeddings. + + Args: + inputs: Input with shape ``[B, S, D]`` or ``[B, S, H, D]``. + + Returns: + Rotated input with the same shape as ``inputs``. + """ + if self.embedding_dims != inputs.shape[-1]: + raise ValueError( + "The embedding dims of the rotary position embedding must " + "match the hidden dimension of the inputs." + ) + + position = torch.arange( + inputs.shape[1], + device=inputs.device, + dtype=torch.float32, + ).unsqueeze(0) + + if inputs.dim() == 4: + position = position.unsqueeze(-1).unsqueeze(-1) + timescale = self.timescale.view(1, 1, 1, -1) + elif inputs.dim() == 3: + position = position.unsqueeze(-1) + timescale = self.timescale.view(1, 1, -1) + else: + raise ValueError("Inputs must be of rank 3 or 4.") + + sinusoid = position.float() / timescale + sin = sinusoid.sin().to(inputs.dtype) + cos = sinusoid.cos().to(inputs.dtype) + first_half, second_half = inputs.chunk(2, dim=-1) + first = first_half * cos - second_half * sin + second = second_half * cos + first_half * sin + return torch.cat((first, second), dim=-1) + + +class MultiHeadAttention(torch.nn.Module): + """Apply TimesFM-3 multi-head attention. + + Args: + num_heads: Number of attention heads. + in_features: Input and output dimension. + use_rotary_position_embeddings: Whether to apply rotary positional + embeddings to queries and keys. + causal_attention: Whether queries can only attend to preceding keys. + use_bias: Whether projection layers use bias parameters. + qk_norm: Query and key normalization. + v_norm: Value normalization. + use_sdpa: Whether to use PyTorch scaled dot-product attention. + rescale_logits: Whether to cancel TimesFM's square-root + head-dimension logit scaling. + device: Device on which to create parameters and buffers. + dtype: Data type of parameters. + """ + + def __init__( + self, + num_heads: int, + in_features: int, + use_rotary_position_embeddings: bool = True, + causal_attention: bool = True, + use_bias: bool = False, + qk_norm: Literal["rms", "none"] = "rms", + v_norm: Literal["rms", "none"] = "none", + use_sdpa: bool = False, + rescale_logits: bool = False, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + factory_kwargs: dict[str, Any] = {"device": device, "dtype": dtype} + + self.num_heads = num_heads + self.in_features = in_features + self.causal_attention = causal_attention + self.head_dim = in_features // num_heads + self.use_sdpa = use_sdpa + self.rescale_logits = rescale_logits + + self.query_proj = torch.nn.Linear( + in_features, + in_features, + bias=use_bias, + **factory_kwargs, + ) + self.key_proj = torch.nn.Linear( + in_features, + in_features, + bias=use_bias, + **factory_kwargs, + ) + self.value_proj = torch.nn.Linear( + in_features, + in_features, + bias=use_bias, + **factory_kwargs, + ) + self.out_proj = torch.nn.Linear( + in_features, + in_features, + bias=use_bias, + **factory_kwargs, + ) + + if qk_norm == "rms": + self.query_ln = torch.nn.RMSNorm( + self.head_dim, + **factory_kwargs, + ) + self.key_ln = torch.nn.RMSNorm( + self.head_dim, + **factory_kwargs, + ) + elif qk_norm == "none": + self.query_ln = None + self.key_ln = None + else: + raise AssertionError(f"Unhandled QK normalization: {qk_norm}") + + if v_norm == "rms": + self.value_ln = torch.nn.RMSNorm( + self.head_dim, + elementwise_affine=False, + **factory_kwargs, + ) + elif v_norm == "none": + self.value_ln = None + else: + raise AssertionError(f"Unhandled value normalization: {v_norm}") + + if use_rotary_position_embeddings: + self.rotary_position_embedding = RotaryPositionalEmbedding( + embedding_dims=self.head_dim, + device=device, + ) + else: + self.rotary_position_embedding = None + + self.per_dim_scale = PerDimScale( + num_dims=self.head_dim, + **factory_kwargs, + ) + + self.register_load_state_dict_post_hook(self._materialize_rope) + + def _materialize_rope( + self, + module: torch.nn.Module, + incompatible_keys: Any, + ) -> None: + del module, incompatible_keys + if self.rotary_position_embedding is not None: + self.rotary_position_embedding._reset_timescale( + self.query_proj.weight.device + ) + + def forward( + self, + inputs_q: Tensor, + *, + patch_mask: Tensor | None = None, + ) -> tuple[Tensor, Tensor]: + """Apply multi-head attention. + + Args: + inputs_q: Inputs with shape ``[B, N, D]``. + patch_mask: Masked patches with shape ``[B, N]``. + + Returns: + Attention output with shape ``[B, N, D]`` and the attention mask. + """ + batch_size, num_patches, _ = inputs_q.shape + if patch_mask is None: + patch_mask = torch.zeros( + batch_size, + num_patches, + dtype=torch.bool, + device=inputs_q.device, + ) + + projection_shape = ( + batch_size, + num_patches, + self.num_heads, + self.head_dim, + ) + query = self.query_proj(inputs_q).view(projection_shape) + key = self.key_proj(inputs_q).view(projection_shape) + value = self.value_proj(inputs_q).view(projection_shape) + + if self.rotary_position_embedding is not None: + query = self.rotary_position_embedding(query) + key = self.rotary_position_embedding(key) + + if self.query_ln is not None: + query = self.query_ln(query) + if self.key_ln is not None: + key = self.key_ln(key) + query = self.per_dim_scale(query) + if self.value_ln is not None: + value = self.value_ln(value) + + attn_mask = make_attn_mask(patch_mask, causal=self.causal_attention) + query = query.transpose(1, 2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + expanded_mask = attn_mask.expand(-1, self.num_heads, -1, -1) + if self.use_sdpa: + scale = 1.0 if self.rescale_logits else math.sqrt(self.head_dim) + x = F.scaled_dot_product_attention( + query, + key, + value, + attn_mask=expanded_mask, + scale=scale, + ) + else: + query = query * math.sqrt(self.head_dim) + logits = query.matmul(key.transpose(-2, -1)) + if self.rescale_logits: + logits = logits / math.sqrt(self.head_dim) + mask_value = max(-1e9, torch.finfo(logits.dtype).min) + mask_bias = torch.zeros_like(logits).masked_fill( + ~expanded_mask, + mask_value, + ) + logits = logits + mask_bias + x = logits.softmax(dim=-1).matmul(value) + + x = ( + x.transpose(1, 2) + .contiguous() + .view(batch_size, num_patches, self.in_features) + ) + return self.out_proj(x), attn_mask + + +class MixingTransformer(torch.nn.Module): + """Apply temporal attention, variate attention, and a feed-forward block. + + Args: + config: Transformer configuration. + use_variate_attention: Whether to attend across variates. + device: Device on which to create parameters and buffers. + dtype: Data type of parameters. + """ + + def __init__( + self, + config: TransformerConfig, + use_variate_attention: bool = True, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + factory_kwargs: dict[str, Any] = {"device": device, "dtype": dtype} + self.config = config + self.use_variate_attention = use_variate_attention + rescale_logits = not config.use_memory_efficient_attention + + self.pre_seq_attn_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.post_seq_attn_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.seq_attn = MultiHeadAttention( + num_heads=config.num_heads, + in_features=config.model_dims, + use_rotary_position_embeddings=config.use_rope_seq, + qk_norm=config.qk_norm, + v_norm=config.v_norm, + causal_attention=config.causal_attention, + use_bias=config.use_bias, + use_sdpa=config.use_sdpa, + rescale_logits=rescale_logits, + **factory_kwargs, + ) + + if use_variate_attention: + self.pre_var_attn_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.post_var_attn_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.var_attn = MultiHeadAttention( + num_heads=config.num_heads, + in_features=config.model_dims, + use_rotary_position_embeddings=config.use_rope_var, + qk_norm=config.qk_norm, + v_norm=config.v_norm, + causal_attention=False, + use_bias=config.use_bias, + use_sdpa=config.use_sdpa, + rescale_logits=rescale_logits, + **factory_kwargs, + ) + + self.pre_ff_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.post_ff_ln = torch.nn.RMSNorm( + config.model_dims, + **factory_kwargs, + ) + self.ff0 = torch.nn.Linear( + config.model_dims, + config.hidden_dims, + bias=config.use_bias, + **factory_kwargs, + ) + self.ff1 = torch.nn.Linear( + config.hidden_dims, + config.model_dims, + bias=config.use_bias, + **factory_kwargs, + ) + self.activation = get_activation_fn(config.ff_activation) + + def forward( + self, + input_embeddings: Tensor, + patch_mask: Tensor, + ) -> tuple[Tensor, Tensor]: + """Apply one mixing transformer layer. + + Args: + input_embeddings: Inputs with shape ``[B, V, N, D]``. + patch_mask: Masked patches with shape ``[B, V, N]``. + + Returns: + Output with shape ``[B, V, N, D]`` and the temporal attention + mask. + """ + batch_size, num_variates, num_patches, model_dims = ( + input_embeddings.shape + ) + seq_input = self.pre_seq_attn_ln(input_embeddings).reshape( + batch_size * num_variates, num_patches, model_dims + ) + seq_patch_mask = patch_mask.reshape( + batch_size * num_variates, num_patches + ) + seq_output, seq_attn_mask = self.seq_attn( + seq_input, patch_mask=seq_patch_mask + ) + seq_output = seq_output.view( + batch_size, num_variates, num_patches, model_dims + ) + hidden = self.post_seq_attn_ln(seq_output) + input_embeddings + + if self.use_variate_attention: + var_input = self.pre_var_attn_ln(hidden) + var_input = var_input.permute(0, 2, 1, 3).reshape( + batch_size * num_patches, num_variates, model_dims + ) + var_patch_mask = patch_mask.permute(0, 2, 1).reshape( + batch_size * num_patches, num_variates + ) + var_output, _ = self.var_attn(var_input, patch_mask=var_patch_mask) + var_output = var_output.view( + batch_size, num_patches, num_variates, model_dims + ).permute(0, 2, 1, 3) + hidden = self.post_var_attn_ln(var_output) + hidden + + ff_output = self.ff1(self.activation(self.ff0(self.pre_ff_ln(hidden)))) + return self.post_ff_ln(ff_output) + hidden, seq_attn_mask + + +class StackedMixingTransformer(torch.nn.Module): + """Apply a stack of TimesFM-3 mixing transformer layers. + + Args: + config: Stacked transformer configuration. + use_variate_attention: Whether to attend across variates. + device: Device on which to create parameters and buffers. + dtype: Data type of parameters. + """ + + def __init__( + self, + config: StackedTransformersConfig, + use_variate_attention: bool = True, + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + ) -> None: + super().__init__() + self.config = config + self.layers = torch.nn.ModuleList( + [ + MixingTransformer( + config=config.transformer, + use_variate_attention=use_variate_attention, + device=device, + dtype=dtype, + ) + for _ in range(config.num_layers) + ] + ) + + def forward( + self, + input_embeddings: Tensor, + patch_mask: Tensor, + ) -> tuple[Tensor, list[Tensor]]: + """Apply the mixing transformer stack. + + Args: + input_embeddings: Inputs with shape ``[B, V, N, D]``. + patch_mask: Masked patches with shape ``[B, V, N]``. + + Returns: + Output with shape ``[B, V, N, D]`` and the temporal attention + mask from every layer. + """ + output = input_embeddings + attn_masks = [] + for layer in self.layers: + output, layer_mask = layer(output, patch_mask) + attn_masks.append(layer_mask) + return output, attn_masks diff --git a/sdm/models/timesfm3/util.py b/sdm/models/timesfm3/util.py index 10fdbfde4..e6bbd6172 100644 --- a/sdm/models/timesfm3/util.py +++ b/sdm/models/timesfm3/util.py @@ -15,7 +15,11 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +from collections.abc import Callable +from typing import Literal + import torch +import torch.nn.functional as F from torch import Tensor _TOLERANCE = 1e-6 @@ -225,6 +229,26 @@ def get_output_patch_via_roll( return result, wrap_mask.unsqueeze(0).unsqueeze(0) +def get_activation_fn( + activation_name: Literal["relu", "swish", "none"], +) -> Callable[[Tensor], Tensor]: + """Return an activation function by name. + + Args: + activation_name: Activation name. + + Returns: + The corresponding tensor operation. + """ + if activation_name == "relu": + return F.relu + if activation_name == "swish": + return F.silu + if activation_name == "none": + return lambda x: x + raise AssertionError(f"Unhandled activation: {activation_name}") + + def stitch_patches( patch_preds: Tensor, patch_len: int, diff --git a/test/models/timesfm3/reference/generate.py b/test/models/timesfm3/reference/generate.py new file mode 100644 index 000000000..47c224593 --- /dev/null +++ b/test/models/timesfm3/reference/generate.py @@ -0,0 +1,169 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Generate TimesFM-3 reference fixtures with pinned Google code.""" + +from __future__ import annotations + +import argparse +import importlib +import json +import subprocess +import sys +from pathlib import Path +from typing import Any + +import torch + +UPSTREAM_REVISION = "e31dadd84cb26bd5153fde6687502b8312e918fb" +WEIGHT_RECIPE = { + "offset_step": 3, + "modulus": 19, + "center": 9, + "divisor": 32, +} +INTERNAL_CONFIG: dict[str, Any] = { + "input_patch_len": 2, + "output_patch_len": 4, + "quantiles": [0.5], + "residual_block_config": { + "hidden_dims": 4, + "output_dims": 4, + "use_bias": False, + "activation": "relu", + }, + "transformer_config": { + "num_layers": 1, + "transformer": { + "model_dims": 4, + "hidden_dims": 6, + "num_heads": 2, + "attention_norm": "rms", + "feedforward_norm": "rms", + "qk_norm": "rms", + "use_rope_seq": True, + "use_rope_var": False, + "use_bias": False, + "ff_activation": "relu", + "deterministic": True, + }, + }, + "use_variate_attention": False, + "use_stitching": True, + "use_linear_detrending": True, + "use_frozen_running_stats": False, +} +FORWARD_INPUTS: dict[str, Any] = { + "values": [[[[1.0, 3.0], [0.0, 0.0], [0.0, 0.0]]]], + "masks": [[[[False, False], [True, True], [True, True]]]], + "patch_is_target": [[[True, True, True]]], + "patch_cpm_masks": [None, [[False, True, True]]], +} + + +def _fill_state(model: torch.nn.Module) -> list[dict[str, Any]]: + state = model.state_dict() + entries = [] + for index, key in enumerate(sorted(state)): + tensor = state[key] + values = ( + ( + torch.arange(tensor.numel()).reshape(tensor.shape) + + WEIGHT_RECIPE["offset_step"] * index + ) + % WEIGHT_RECIPE["modulus"] + - WEIGHT_RECIPE["center"] + ) / WEIGHT_RECIPE["divisor"] + state[key] = values.to(dtype=tensor.dtype) + entries.append({"key": key, "shape": list(tensor.shape)}) + model.load_state_dict(state, strict=True) + return entries + + +def _upstream_revision(checkout: Path) -> str: + result = subprocess.run( + ["git", "-C", str(checkout), "rev-parse", "HEAD"], + check=True, + capture_output=True, + text=True, + ) + return result.stdout.strip() + + +def _internal_fixture(upstream: Any) -> dict[str, Any]: + model = upstream.TimesFM3Torch(**INTERNAL_CONFIG).eval() + state = _fill_state(model) + values = torch.tensor(FORWARD_INPUTS["values"]) + masks = torch.tensor(FORWARD_INPUTS["masks"]) + patch_is_target = torch.tensor(FORWARD_INPUTS["patch_is_target"]) + forward_outputs = [] + with torch.inference_mode(): + for patch_cpm_mask in FORWARD_INPUTS["patch_cpm_masks"]: + cpm_mask = ( + None + if patch_cpm_mask is None + else torch.tensor(patch_cpm_mask) + ) + logits = model( + { + "values": values, + "masks": masks, + "patch_is_target": patch_is_target, + }, + patch_cpm_mask=cpm_mask, + )["logits"] + forward_outputs.append(logits.tolist()) + + return { + "config": INTERNAL_CONFIG, + "state": state, + "forward": { + "inputs": FORWARD_INPUTS, + "outputs": forward_outputs, + }, + } + + +def generate(checkout: Path) -> dict[str, Any]: + """Generate the fixtures using an exact Google TimesFM checkout.""" + revision = _upstream_revision(checkout) + if revision != UPSTREAM_REVISION: + raise RuntimeError( + f"Expected upstream revision {UPSTREAM_REVISION}, got {revision}." + ) + + sys.path.insert(0, str(checkout / "src")) + upstream = importlib.import_module("timesfm3.torch.model") + return { + "upstream": { + "repository": "https://github.com/google-research/timesfm", + "revision": revision, + }, + "weight_recipe": WEIGHT_RECIPE, + "internal": _internal_fixture(upstream), + } + + +def main() -> None: + """Generate and write the reference fixtures.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "checkout", + type=Path, + help="Path to the pinned google-research/timesfm checkout.", + ) + parser.add_argument( + "--output", + type=Path, + default=Path(__file__).with_name("golden.json"), + ) + args = parser.parse_args() + + fixture = generate(args.checkout.resolve()) + args.output.write_text( + json.dumps(fixture, indent=2, allow_nan=False) + "\n" + ) + + +if __name__ == "__main__": + main() diff --git a/test/models/timesfm3/reference/golden.json b/test/models/timesfm3/reference/golden.json new file mode 100644 index 000000000..aca0a5cab --- /dev/null +++ b/test/models/timesfm3/reference/golden.json @@ -0,0 +1,324 @@ +{ + "upstream": { + "repository": "https://github.com/google-research/timesfm", + "revision": "e31dadd84cb26bd5153fde6687502b8312e918fb" + }, + "weight_recipe": { + "offset_step": 3, + "modulus": 19, + "center": 9, + "divisor": 32 + }, + "internal": { + "config": { + "input_patch_len": 2, + "output_patch_len": 4, + "quantiles": [ + 0.5 + ], + "residual_block_config": { + "hidden_dims": 4, + "output_dims": 4, + "use_bias": false, + "activation": "relu" + }, + "transformer_config": { + "num_layers": 1, + "transformer": { + "model_dims": 4, + "hidden_dims": 6, + "num_heads": 2, + "attention_norm": "rms", + "feedforward_norm": "rms", + "qk_norm": "rms", + "use_rope_seq": true, + "use_rope_var": false, + "use_bias": false, + "ff_activation": "relu", + "deterministic": true + } + }, + "use_variate_attention": false, + "use_stitching": true, + "use_linear_detrending": true, + "use_frozen_running_stats": false + }, + "state": [ + { + "key": "output_head.bias", + "shape": [ + 4 + ] + }, + { + "key": "output_head.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "pre_transformer_resblock.hidden_layer.weight", + "shape": [ + 4, + 12 + ] + }, + { + "key": "pre_transformer_resblock.output_layer.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "pre_transformer_resblock.residual_layer.weight", + "shape": [ + 4, + 12 + ] + }, + { + "key": "transformer_stack.layers.0.ff0.weight", + "shape": [ + 6, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.ff1.weight", + "shape": [ + 4, + 6 + ] + }, + { + "key": "transformer_stack.layers.0.post_ff_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.post_seq_attn_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.pre_ff_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.pre_seq_attn_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.key_ln.weight", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.key_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.out_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.per_dim_scale.per_dim_scale", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.query_ln.weight", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.query_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.value_proj.weight", + "shape": [ + 4, + 4 + ] + } + ], + "forward": { + "inputs": { + "values": [ + [ + [ + [ + 1, + 3 + ], + [ + 0, + 0 + ], + [ + 0, + 0 + ] + ] + ] + ], + "masks": [ + [ + [ + [ + false, + false + ], + [ + true, + true + ], + [ + true, + true + ] + ] + ] + ], + "patch_is_target": [ + [ + [ + true, + true, + true + ] + ] + ], + "patch_cpm_masks": [ + null, + [ + [ + false, + true, + true + ] + ] + ] + }, + "outputs": [ + [ + [ + [ + [ + [ + 1.7138645648956299 + ], + [ + 1.7143478393554688 + ], + [ + 1.7148311138153076 + ], + [ + 1.715314269065857 + ] + ], + [ + [ + 1.692699670791626 + ], + [ + 1.7234890460968018 + ], + [ + 1.754278302192688 + ], + [ + 1.7850675582885742 + ] + ], + [ + [ + 1.6920182704925537 + ], + [ + 1.7233555316925049 + ], + [ + 1.754692792892456 + ], + [ + 1.7860300540924072 + ] + ] + ] + ] + ], + [ + [ + [ + [ + [ + 1.7138645648956299 + ], + [ + 1.7143478393554688 + ], + [ + 1.7148311138153076 + ], + [ + 1.715314269065857 + ] + ], + [ + [ + 1.635363221168518 + ], + [ + 1.657575011253357 + ], + [ + 1.6797866821289062 + ], + [ + 1.7019984722137451 + ] + ], + [ + [ + 1.6271485090255737 + ], + [ + 1.6457258462905884 + ], + [ + 1.6643033027648926 + ], + [ + 1.6828805208206177 + ] + ] + ] + ] + ] + ] + } + } +} diff --git a/test/models/timesfm3/test_core.py b/test/models/timesfm3/test_core.py new file mode 100644 index 000000000..33551a91f --- /dev/null +++ b/test/models/timesfm3/test_core.py @@ -0,0 +1,528 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import json +from pathlib import Path +from typing import Any, cast + +import pytest +import torch + +from sdm.models.timesfm3.configs import ( + ResidualBlockConfig, + StackedTransformersConfig, + TransformerConfig, +) +from sdm.models.timesfm3.core import _TimesFM3Model +from sdm.models.timesfm3.util import get_running_stats +from sdm.testing import withCUDA + + +def _residual_config(output_dims: int = 8) -> ResidualBlockConfig: + return ResidualBlockConfig( + hidden_dims=8, + output_dims=output_dims, + use_bias=False, + activation="relu", + ) + + +def _transformer_config(model_dims: int = 8) -> StackedTransformersConfig: + return StackedTransformersConfig( + num_layers=2, + transformer=TransformerConfig( + model_dims=model_dims, + hidden_dims=12, + num_heads=2, + qk_norm="rms", + use_rope_seq=True, + use_rope_var=False, + use_bias=False, + ff_activation="relu", + ), + ) + + +def _internal_model( + device: torch.device | str, + dtype: torch.dtype | None = None, + *, + use_iterative_cpm_revin: bool = True, + use_linear_detrending: bool = True, +) -> _TimesFM3Model: + return _TimesFM3Model( + input_patch_len=2, + output_patch_len=4, + quantiles=[0.1, 0.5, 0.9], + residual_block_config=_residual_config(), + transformer_config=_transformer_config(), + use_iterative_cpm_revin=use_iterative_cpm_revin, + use_linear_detrending=use_linear_detrending, + device=device, + dtype=dtype, + ) + + +@withCUDA +def test_internal_model_assembles_checkpoint_parameters( + device: torch.device, +) -> None: + source = _internal_model(device) + state = source.state_dict() + assert state["pre_transformer_resblock.hidden_layer.weight"].shape == ( + 8, + 12, + ) + assert state[ + "transformer_stack.layers.0.seq_attn.query_proj.weight" + ].shape == ( + 8, + 8, + ) + assert state["output_head.weight"].shape == (12, 8) + + loaded = _internal_model("meta") + loaded.load_state_dict(state, strict=True, assign=True) + assert all(parameter.device == device for parameter in loaded.parameters()) + torch.testing.assert_close( + loaded.output_head.weight, source.output_head.weight + ) + + +@withCUDA +def test_internal_model_preprocesses_patches(device: torch.device) -> None: + model = _internal_model(device).eval() + values = torch.tensor( + [[[[1.0, 3.0], [10.0, 14.0], [100.0, 200.0]]]], + device=device, + ) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.ones(1, 1, 3, dtype=torch.bool, device=device) + + residual_input, transformer_input, patch_mask, stats, counts = ( + model._preprocess( + values, + masks, + patch_is_target, + freeze_after=0, + ) + ) + + running_mean, running_std = stats + assert residual_input.shape == (1, 1, 3, 12) + assert transformer_input.shape == (1, 1, 3, 8) + assert patch_mask.shape == (1, 1, 3) + torch.testing.assert_close( + running_mean, + torch.tensor([[[2.0, 2.0, 2.0]]], device=device), + ) + torch.testing.assert_close( + running_std, + torch.tensor([[[1.0, 1.0, 1.0]]], device=device), + ) + torch.testing.assert_close( + counts, + torch.tensor([[[2.0, 4.0, 6.0]]], device=device), + ) + assert residual_input[..., 8:].bool().all() + + +@withCUDA +def test_internal_model_preserves_covariates_under_cpm( + device: torch.device, +) -> None: + model = _internal_model(device).eval() + values = torch.tensor( + [ + [ + [[0.0, 2.0], [4.0, 6.0], [8.0, 10.0]], + [[10.0, 12.0], [10.0, 12.0], [10.0, 12.0]], + ] + ], + device=device, + ) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.tensor([[[True] * 3, [False] * 3]], device=device) + patch_cpm_mask = torch.tensor([[False, True, False]], device=device) + + residual_input, _, patch_mask, stats, counts = model._preprocess( + values, + masks, + patch_is_target, + patch_cpm_mask=patch_cpm_mask, + ) + running_mean, _ = stats + + assert running_mean[0, 0, 1] == 3.0 + assert counts[0, 0, 1] == 4.0 + # Current values, future values, current masks, then future masks. + assert not residual_input[0, 0, 1, :6].bool().any() + torch.testing.assert_close( + residual_input[0, 1, 0, 2:6], + torch.tensor([-1.0, 1.0, -1.0, 1.0], device=device), + ) + torch.testing.assert_close( + residual_input[0, 1, 1, 2:6], + torch.tensor([-1.0, 1.0, 0.0, 0.0], device=device), + ) + assert residual_input[0, 0, :, 8:12].bool().all() + assert not residual_input[0, 1, 1, 8:10].bool().any() + assert residual_input[0, 1, 1, 10:12].bool().all() + assert residual_input[0, 1, 2, 8:12].bool().all() + expected_patch_mask = torch.tensor( + [[[False, True, False], [False, False, False]]], device=device + ) + torch.testing.assert_close(patch_mask, expected_patch_mask) + + +@withCUDA +def test_internal_model_forward(device: torch.device) -> None: + model = _internal_model(device).eval() + values = torch.arange( + 2 * 3 * 4 * 2, + device=device, + dtype=torch.float32, + ).reshape(2, 3, 4, 2) + masks = torch.zeros_like(values, dtype=torch.bool) + masks[:, :, 0] = True + masks[:, :, 2] = True + patch_is_target = torch.ones(2, 3, 4, dtype=torch.bool, device=device) + + with torch.inference_mode(): + outputs = model( + values, + masks, + patch_is_target, + return_aux_outputs=True, + ) + + assert outputs["logits"].shape == (2, 3, 4, 4, 3) + assert torch.isfinite(outputs["logits"]).all() + assert outputs["__call__:resblock_input"].shape == (2, 3, 4, 12) + assert outputs["__call__:transformer_input"].shape == (2, 3, 4, 8) + assert outputs["__call__:transformer_output"].shape == (2, 3, 4, 8) + + attention_masks = outputs["__call__:seq_attn_mask"] + assert len(attention_masks) == 2 + expected_mask = torch.tensor( + [ + [False, False, False, False], + [False, True, False, False], + [False, True, True, False], + [False, True, True, True], + ], + device=device, + ) + torch.testing.assert_close(attention_masks[0][0, 0], expected_mask) + + +@withCUDA +def test_internal_model_denormalizes_output_head( + device: torch.device, +) -> None: + model = _internal_model( + device, + use_iterative_cpm_revin=False, + ).eval() + normalized_logits = ( + torch.arange( + 12, + device=device, + dtype=torch.float32, + ).reshape(4, 3) + / 10.0 + ) + with torch.no_grad(): + for parameter in model.parameters(): + parameter.zero_() + model.output_head.bias.copy_(normalized_logits.flatten()) + + values = torch.tensor( + [[[[1.0, 3.0], [5.0, 7.0], [9.0, 11.0]]]], + device=device, + ) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.ones(1, 1, 3, dtype=torch.bool, device=device) + + with torch.inference_mode(): + logits = model(values, masks, patch_is_target)["logits"] + + _, running_mean, running_std = get_running_stats(values, masks) + expected = ( + normalized_logits[None, None, None] * running_std[..., None, None] + + running_mean[..., None, None] + ) + torch.testing.assert_close(logits, expected) + + +@withCUDA +def test_internal_model_bfloat16(device: torch.device) -> None: + model = _internal_model(device, dtype=torch.bfloat16).eval() + values = torch.arange( + 2 * 2 * 3 * 2, + device=device, + dtype=torch.float32, + ).reshape(2, 2, 3, 2) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.ones(2, 2, 3, dtype=torch.bool, device=device) + + with torch.inference_mode(): + outputs = model( + values, + masks, + patch_is_target, + return_aux_outputs=True, + ) + + assert outputs["__call__:resblock_input"].dtype == torch.bfloat16 + assert outputs["__call__:transformer_output"].dtype == torch.bfloat16 + assert torch.isfinite(outputs["logits"]).all() + + +@withCUDA +def test_internal_model_loads_from_meta(device: torch.device) -> None: + expected = _internal_model(device).eval() + model = _internal_model("meta").eval() + assert all( + parameter.device.type == "meta" for parameter in model.parameters() + ) + + model.load_state_dict(expected.state_dict(), assign=True) + model.eval() + values = torch.arange( + 2 * 2 * 3 * 2, + device=device, + dtype=torch.float32, + ).reshape(2, 2, 3, 2) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.ones(2, 2, 3, dtype=torch.bool, device=device) + + with torch.inference_mode(): + actual_logits = model(values, masks, patch_is_target)["logits"] + expected_logits = expected(values, masks, patch_is_target)["logits"] + + assert all(parameter.device == device for parameter in model.parameters()) + assert all(buffer.device == device for buffer in model.buffers()) + torch.testing.assert_close(actual_logits, expected_logits) + + +def test_internal_model_rejects_incompatible_configuration() -> None: + with pytest.raises(ValueError, match="must be a multiple"): + _TimesFM3Model( + input_patch_len=2, + output_patch_len=3, + residual_block_config=_residual_config(), + transformer_config=_transformer_config(), + use_stitching=False, + ) + + with pytest.raises(ValueError, match="dimensions must match"): + _TimesFM3Model( + input_patch_len=2, + output_patch_len=4, + residual_block_config=_residual_config(output_dims=4), + transformer_config=_transformer_config(model_dims=8), + ) + + with pytest.raises(ValueError, match="Stitching requires"): + _TimesFM3Model( + input_patch_len=2, + output_patch_len=2, + residual_block_config=_residual_config(), + transformer_config=_transformer_config(), + ) + + model = _internal_model("cpu") + with pytest.raises(ValueError, match="does not match"): + model( + torch.zeros(1, 1, 1, 3), + torch.zeros(1, 1, 1, 3, dtype=torch.bool), + torch.ones(1, 1, 1, dtype=torch.bool), + ) + + +@withCUDA +def test_internal_model_sanitizes_inputs(device: torch.device) -> None: + model = _TimesFM3Model( + input_patch_len=2, + output_patch_len=4, + quantiles=[0.1, 0.5, 0.9], + residual_block_config=_residual_config(), + transformer_config=_transformer_config(), + value_clip=5.0, + device=device, + ).eval() + values = torch.tensor( + [[[[float("nan"), float("inf")], [float("-inf"), 10.0]]]], + device=device, + ) + sanitized = torch.tensor( + [[[[0.0, 5.0], [-5.0, 5.0]]]], + device=device, + ) + masks = torch.zeros_like(values, dtype=torch.bool) + patch_is_target = torch.ones(1, 1, 2, dtype=torch.bool, device=device) + + with torch.inference_mode(): + actual = model( + values, + masks, + patch_is_target, + return_aux_outputs=True, + ) + expected = model( + sanitized, + masks, + patch_is_target, + return_aux_outputs=True, + ) + + torch.testing.assert_close(actual["logits"], expected["logits"]) + torch.testing.assert_close( + actual["__call__:resblock_input"], + expected["__call__:resblock_input"], + ) + + +@withCUDA +def test_internal_model_applies_cpm_revin_refinement( + device: torch.device, +) -> None: + refined_model = _internal_model( + device, + use_iterative_cpm_revin=True, + ).eval() + frozen_model = _internal_model( + device, + use_iterative_cpm_revin=False, + ).eval() + normalized_logits = ( + torch.arange( + 12, + device=device, + dtype=torch.float32, + ) + / 10.0 + ) + with torch.no_grad(): + for parameter in refined_model.parameters(): + parameter.zero_() + refined_model.output_head.bias.copy_(normalized_logits) + frozen_model.load_state_dict(refined_model.state_dict()) + + values = torch.tensor( + [[[[1.0, 3.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]]], + device=device, + ) + masks = torch.tensor( + [[[[False, False], [True, True], [True, True], [True, True]]]], + device=device, + ) + patch_is_target = torch.ones(1, 1, 4, dtype=torch.bool, device=device) + patch_cpm_mask = torch.tensor( + [[False, True, True, True]], + device=device, + ) + + with torch.inference_mode(): + refined = refined_model( + values, + masks, + patch_is_target, + patch_cpm_mask=patch_cpm_mask, + )["logits"] + frozen = frozen_model( + values, + masks, + patch_is_target, + patch_cpm_mask=patch_cpm_mask, + )["logits"] + + torch.testing.assert_close(refined[:, :, 0], frozen[:, :, 0]) + assert not torch.equal(refined[:, :, 1:], frozen[:, :, 1:]) + + +def _reference_fixture() -> dict[str, Any]: + path = Path(__file__).with_name("reference") / "golden.json" + return cast(dict[str, Any], json.loads(path.read_text())) + + +def _model_config(config: dict[str, Any]) -> dict[str, Any]: + model_config = dict(config) + stack_config = dict(model_config["transformer_config"]) + transformer_config = dict(stack_config["transformer"]) + for name in ("attention_norm", "feedforward_norm", "deterministic"): + transformer_config.pop(name) + stack_config["transformer"] = transformer_config + model_config["transformer_config"] = stack_config + return model_config + + +def _load_reference_weights( + model: _TimesFM3Model, + model_fixture: dict[str, Any], + recipe: dict[str, int], +) -> None: + state = model.state_dict() + actual_state = [ + {"key": key, "shape": list(tensor.shape)} + for key, tensor in sorted(state.items()) + ] + assert actual_state == model_fixture["state"] + + for index, key in enumerate(sorted(state)): + tensor = state[key] + values = ( + ( + torch.arange(tensor.numel(), device=tensor.device).reshape( + tensor.shape + ) + + recipe["offset_step"] * index + ) + % recipe["modulus"] + - recipe["center"] + ) / recipe["divisor"] + state[key] = values.to(dtype=tensor.dtype) + model.load_state_dict(state, strict=True) + + +def _golden_model(device: torch.device) -> _TimesFM3Model: + fixture = _reference_fixture() + model_fixture = fixture["internal"] + model = _TimesFM3Model( + **_model_config(model_fixture["config"]), + device=device, + ).eval() + _load_reference_weights(model, model_fixture, fixture["weight_recipe"]) + return model + + +@pytest.mark.parametrize("case_index", [0, 1]) +@withCUDA +def test_internal_model_matches_pinned_upstream( + device: torch.device, + case_index: int, +) -> None: + fixture = _reference_fixture()["internal"]["forward"] + inputs = fixture["inputs"] + model = _golden_model(device) + values = torch.tensor(inputs["values"], device=device) + masks = torch.tensor(inputs["masks"], device=device) + patch_is_target = torch.tensor(inputs["patch_is_target"], device=device) + patch_cpm_mask = inputs["patch_cpm_masks"][case_index] + cpm_mask = ( + None + if patch_cpm_mask is None + else torch.tensor(patch_cpm_mask, device=device) + ) + + with torch.inference_mode(): + actual = model( + values, + masks, + patch_is_target, + patch_cpm_mask=cpm_mask, + )["logits"] + + expected = actual.new_tensor(fixture["outputs"][case_index]) + torch.testing.assert_close(actual, expected) diff --git a/test/models/timesfm3/test_cpm_revin_refine.py b/test/models/timesfm3/test_cpm_revin_refine.py new file mode 100644 index 000000000..86eb359a8 --- /dev/null +++ b/test/models/timesfm3/test_cpm_revin_refine.py @@ -0,0 +1,250 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import torch + +from sdm.models.timesfm3.cpm_revin_refine import ( + cpm_iterative_revin_refine, +) +from sdm.testing import withCUDA + + +def _raw_logits( + median_logits: torch.Tensor, + num_quantiles: int, + median_q_idx: int, +) -> torch.Tensor: + *leading, rolls, patch_len = median_logits.shape + logits = median_logits.new_zeros( + *leading, + rolls, + patch_len, + num_quantiles, + ) + logits[..., median_q_idx] = median_logits + return logits.flatten(start_dim=-3) + + +@withCUDA +def test_cpm_revin_refine_without_mask_is_identity( + device: torch.device, +) -> None: + batch_size, num_variates, num_patches = 2, 3, 5 + rolls, patch_len, num_quantiles = 2, 4, 3 + output_dims = rolls * patch_len * num_quantiles + raw_logits = torch.arange( + batch_size * num_variates * num_patches * output_dims, + device=device, + dtype=torch.float32, + ).reshape(batch_size, num_variates, num_patches, output_dims) + revin_n = torch.full( + (batch_size, num_variates, num_patches), + 4.0, + device=device, + ) + revin_mu = torch.arange( + batch_size * num_variates * num_patches, + device=device, + dtype=torch.float32, + ).reshape(batch_size, num_variates, num_patches) + revin_sigma = revin_mu + 1.0 + patch_cpm_mask = torch.zeros( + batch_size, + num_patches, + dtype=torch.bool, + device=device, + ) + + refined_mu, refined_sigma = cpm_iterative_revin_refine( + raw_logits=raw_logits, + revin_n=revin_n, + revin_mu=revin_mu, + revin_sigma=revin_sigma, + patch_cpm_mask=patch_cpm_mask, + median_q_idx=1, + rolls=rolls, + patch_len=patch_len, + num_quantiles=num_quantiles, + ) + + torch.testing.assert_close(refined_mu, revin_mu) + torch.testing.assert_close(refined_sigma, revin_sigma) + + +@withCUDA +def test_cpm_revin_refine_updates_from_anchor_predictions( + device: torch.device, +) -> None: + median_logits = torch.tensor( + [ + [ + [ + [[2.0, 4.0], [6.0, 8.0]], + [[0.0, 0.0], [0.0, 0.0]], + [[0.0, 0.0], [0.0, 0.0]], + [[1.0, 1.0], [2.0, 2.0]], + [[0.0, 0.0], [0.0, 0.0]], + ] + ] + ], + device=device, + ) + raw_logits = _raw_logits( + median_logits, + num_quantiles=3, + median_q_idx=1, + ) + revin_n = torch.tensor( + [[[2.0, 2.0, 2.0, 1.0, 1.0]]], + device=device, + ) + revin_mu = torch.tensor( + [[[0.0, 20.0, 30.0, 10.0, 40.0]]], + device=device, + ) + revin_sigma = torch.tensor( + [[[1.0, 2.0, 3.0, 2.0, 4.0]]], + device=device, + ) + patch_cpm_mask = torch.tensor( + [[False, True, True, False, True]], + device=device, + ) + + refined_mu, refined_sigma = cpm_iterative_revin_refine( + raw_logits=raw_logits, + revin_n=revin_n, + revin_mu=revin_mu, + revin_sigma=revin_sigma, + patch_cpm_mask=patch_cpm_mask, + median_q_idx=1, + rolls=2, + patch_len=2, + num_quantiles=3, + ) + + expected_mu = torch.tensor( + [[[0.0, 1.5, 10.0 / 3.0, 10.0, 34.0 / 3.0]]], + device=device, + ) + expected_sigma = torch.tensor( + [ + [ + [ + 1.0, + 3.25**0.5, + (83.0 / 9.0) ** 0.5, + 2.0, + (20.0 / 9.0) ** 0.5, + ] + ] + ], + device=device, + ) + torch.testing.assert_close(refined_mu, expected_mu) + torch.testing.assert_close(refined_sigma, expected_sigma) + + +@withCUDA +def test_cpm_revin_refine_clips_anchor_predictions( + device: torch.device, +) -> None: + median_logits = torch.tensor( + [[[[[100.0]], [[0.0]]]]], + device=device, + ) + raw_logits = _raw_logits( + median_logits, + num_quantiles=1, + median_q_idx=0, + ) + revin_n = torch.ones(1, 1, 2, device=device) + revin_mu = torch.zeros(1, 1, 2, device=device) + revin_sigma = torch.ones(1, 1, 2, device=device) + patch_cpm_mask = torch.tensor([[False, True]], device=device) + + refined_mu, refined_sigma = cpm_iterative_revin_refine( + raw_logits=raw_logits, + revin_n=revin_n, + revin_mu=revin_mu, + revin_sigma=revin_sigma, + patch_cpm_mask=patch_cpm_mask, + median_q_idx=0, + rolls=1, + patch_len=1, + num_quantiles=1, + value_clip=5.0, + ) + + expected_mu = torch.tensor([[[0.0, 2.5]]], device=device) + expected_sigma = torch.tensor( + [[[1.0, 6.75**0.5]]], + device=device, + ) + torch.testing.assert_close(refined_mu, expected_mu) + torch.testing.assert_close(refined_sigma, expected_sigma) + + +@withCUDA +def test_cpm_revin_refine_tracks_each_batch_independently( + device: torch.device, +) -> None: + batch_size, num_patches = 2, 6 + rolls, patch_len, num_quantiles = 2, 2, 3 + output_dims = rolls * patch_len * num_quantiles + raw_logits = torch.arange( + batch_size * num_patches * output_dims, + device=device, + dtype=torch.float32, + ).reshape(batch_size, 1, num_patches, output_dims) + raw_logits = raw_logits / 10.0 + revin_n = torch.arange( + 1, + num_patches + 1, + device=device, + dtype=torch.float32, + ).reshape(1, 1, num_patches) + revin_n = revin_n.expand(batch_size, -1, -1) + revin_mu = torch.tensor( + [[[0.0, 1.0, 2.0, 3.0, 4.0, 5.0]]], + device=device, + ).expand(batch_size, -1, -1) + revin_sigma = revin_mu + 1.0 + patch_cpm_mask = torch.tensor( + [ + [False, True, True, False, True, True], + [False, False, True, True, True, False], + ], + device=device, + ) + kwargs = { + "median_q_idx": 1, + "rolls": rolls, + "patch_len": patch_len, + "num_quantiles": num_quantiles, + } + + batched = cpm_iterative_revin_refine( + raw_logits=raw_logits, + revin_n=revin_n, + revin_mu=revin_mu, + revin_sigma=revin_sigma, + patch_cpm_mask=patch_cpm_mask, + **kwargs, + ) + independent = [ + cpm_iterative_revin_refine( + raw_logits=raw_logits[index : index + 1], + revin_n=revin_n[index : index + 1], + revin_mu=revin_mu[index : index + 1], + revin_sigma=revin_sigma[index : index + 1], + patch_cpm_mask=patch_cpm_mask[index : index + 1], + **kwargs, + ) + for index in range(batch_size) + ] + + expected_mu = torch.cat([result[0] for result in independent]) + expected_sigma = torch.cat([result[1] for result in independent]) + torch.testing.assert_close(batched[0], expected_mu) + torch.testing.assert_close(batched[1], expected_sigma) diff --git a/test/models/timesfm3/test_dense.py b/test/models/timesfm3/test_dense.py new file mode 100644 index 000000000..6591f414d --- /dev/null +++ b/test/models/timesfm3/test_dense.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch + +from sdm.models.timesfm3.configs import ResidualBlockConfig +from sdm.models.timesfm3.dense import ResidualBlock +from sdm.testing import withCUDA + + +@withCUDA +def test_residual_block_loads_checkpoint_weights(device: torch.device) -> None: + config = ResidualBlockConfig( + hidden_dims=2, + output_dims=2, + use_bias=False, + activation="relu", + ) + block = ResidualBlock(config, input_dims=3, device=device) + block.load_state_dict( + { + "hidden_layer.weight": torch.tensor( + [[1.0, -1.0, 0.0], [0.0, 1.0, 1.0]], + device=device, + ), + "output_layer.weight": torch.tensor( + [[1.0, 2.0], [-1.0, 1.0]], + device=device, + ), + "residual_layer.weight": torch.tensor( + [[1.0, 0.0, 0.0], [0.0, 0.0, 1.0]], + device=device, + ), + } + ) + + x = torch.tensor([[2.0, 1.0, 3.0]], device=device) + + output = block(x) + + torch.testing.assert_close(output, output.new_tensor([[11.0, 6.0]])) + + +@withCUDA +def test_residual_block_explicitly_sets_input_dimension( + device: torch.device, +) -> None: + config = ResidualBlockConfig( + hidden_dims=3, + output_dims=4, + use_bias=True, + activation="swish", + ) + block = ResidualBlock(config, input_dims=5).to( + device=device, dtype=torch.float64 + ) + parameter_ids = tuple(map(id, block.parameters())) + x = torch.randn(2, 5, device=device, dtype=torch.float64) + + output = block(x) + assert tuple(map(id, block.parameters())) == parameter_ids + + assert output.shape == (2, 4) + assert output.device == device + assert output.dtype == x.dtype + + +def test_residual_block_meta_device() -> None: + config = ResidualBlockConfig( + hidden_dims=3, + output_dims=4, + use_bias=True, + activation="swish", + prenorm="rms", + ) + block = ResidualBlock(config, input_dims=5, device="meta") + assert all( + parameter.device.type == "meta" for parameter in block.parameters() + ) + + +@withCUDA +def test_residual_block_identity_skip(device: torch.device) -> None: + config = ResidualBlockConfig( + hidden_dims=3, + output_dims=2, + use_bias=False, + activation="none", + identity_skip=True, + ) + block = ResidualBlock(config, input_dims=2, device=device) + with torch.no_grad(): + block.hidden_layer.weight.zero_() + block.output_layer.weight.zero_() + x = torch.tensor([[1.0, -2.0]], device=device) + + output = block(x) + + torch.testing.assert_close(output, x) + + +def test_residual_block_identity_skip_requires_matching_dimensions() -> None: + config = ResidualBlockConfig( + hidden_dims=3, + output_dims=1, + use_bias=False, + activation="none", + identity_skip=True, + ) + + with pytest.raises(ValueError, match="identity_skip requires"): + ResidualBlock(config, input_dims=4) + + +@withCUDA +def test_residual_block_rms_prenorm(device: torch.device) -> None: + config = ResidualBlockConfig( + hidden_dims=2, + output_dims=2, + use_bias=False, + activation="none", + identity_skip=True, + prenorm="rms", + ) + block = ResidualBlock(config, input_dims=2, device=device) + with torch.no_grad(): + block.hidden_layer.weight.copy_(torch.eye(2, device=device)) + block.output_layer.weight.copy_(torch.eye(2, device=device)) + x = torch.tensor([[3.0, 4.0]], device=device) + + output = block(x) + + expected = x + torch.nn.functional.rms_norm(x, (2,)) + torch.testing.assert_close(output, expected) diff --git a/test/models/timesfm3/test_primitives.py b/test/models/timesfm3/test_primitives.py new file mode 100644 index 000000000..fc8dc19c8 --- /dev/null +++ b/test/models/timesfm3/test_primitives.py @@ -0,0 +1,38 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import torch + +from sdm.models.timesfm3.normalization import PerDimScale +from sdm.testing import withCUDA + + +@withCUDA +def test_per_dim_scale(device: torch.device) -> None: + num_dims = 8 + module = PerDimScale(num_dims=num_dims, device=device) + with torch.no_grad(): + module.per_dim_scale.fill_(1.0) + tensor = torch.ones(2, 3, num_dims, device=device) + + out = module(tensor) + + torch.testing.assert_close( + out, + torch.full_like(tensor, 0.669855025622358), + ) + + +@withCUDA +def test_per_dim_scale_dtype_device(device: torch.device) -> None: + dtype = torch.float64 + module = PerDimScale(num_dims=4, device=device, dtype=dtype) + tensor = torch.ones(2, 4, device=device, dtype=dtype) + + out = module(tensor) + + assert out.dtype == dtype + assert out.device == device + assert module.per_dim_scale.dtype == dtype + assert module.per_dim_scale.device == device + assert tuple(module.state_dict()) == ("per_dim_scale",) diff --git a/test/models/timesfm3/test_transformer.py b/test/models/timesfm3/test_transformer.py new file mode 100644 index 000000000..b2c515d4a --- /dev/null +++ b/test/models/timesfm3/test_transformer.py @@ -0,0 +1,550 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import math +from dataclasses import replace + +import pytest +import torch + +from sdm.models.timesfm3.configs import ( + StackedTransformersConfig, + TransformerConfig, +) +from sdm.models.timesfm3.transformer import ( + MixingTransformer, + MultiHeadAttention, + RotaryPositionalEmbedding, + StackedMixingTransformer, + make_attn_mask, +) +from sdm.testing import withCUDA + + +@withCUDA +def test_make_attn_mask(device: torch.device) -> None: + patch_mask = torch.tensor( + [[True, False, False], [False, True, False]], device=device + ) + mask = make_attn_mask(patch_mask) + causal = torch.ones(3, 3, dtype=torch.bool, device=device).tril() + expected = causal[None, None] & ~patch_mask[:, None, None, :] + torch.testing.assert_close(mask, expected) + + +@withCUDA +def test_make_attn_mask_noncausal(device: torch.device) -> None: + patch_mask = torch.tensor([[True, False, False]], device=device) + mask = make_attn_mask(patch_mask, causal=False) + torch.testing.assert_close(mask, ~patch_mask[:, None, None, :]) + + +@withCUDA +@pytest.mark.parametrize("rank", [3, 4]) +def test_rotary_positional_embedding( + device: torch.device, + rank: int, +) -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4, device=device) + shape = (2, 3, 4) if rank == 3 else (2, 3, 2, 4) + inputs = torch.arange( + math.prod(shape), + device=device, + dtype=torch.float32, + ).reshape(shape) + output = rope(inputs) + + torch.testing.assert_close( + output.square().sum(dim=-1), + inputs.square().sum(dim=-1), + ) + torch.testing.assert_close(output[0, 0], inputs[0, 0]) + assert output.shape == inputs.shape + assert output.device == device + + +def test_rotary_positional_embedding_excludes_buffer_from_checkpoint() -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4) + + assert rope.state_dict() == {} + + +def test_rotary_positional_embedding_rejects_wrong_dimension() -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4) + + with pytest.raises(ValueError, match="must match the hidden dimension"): + rope(torch.zeros(1, 2, 6)) + + +def test_rotary_positional_embedding_rejects_wrong_rank() -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4) + + with pytest.raises(ValueError, match="rank 3 or 4"): + rope(torch.zeros(2, 4)) + + +def test_rotary_positional_embedding_meta_device() -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4, device="meta") + + assert rope.timescale.device.type == "meta" + + +@withCUDA +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_rotary_positional_embedding_rotation( + device: torch.device, + dtype: torch.dtype, +) -> None: + rope = RotaryPositionalEmbedding(embedding_dims=4, device=device) + inputs = torch.tensor( + [[[0.0, 0.0, 0.0, 0.0], [1.0, 1.0, 0.0, 0.0]]], + device=device, + dtype=dtype, + ) + + output = rope(inputs) + + expected = torch.tensor( + [math.cos(1.0), math.cos(0.01), math.sin(1.0), math.sin(0.01)], + device=device, + dtype=dtype, + ) + torch.testing.assert_close(output[0, 1], expected) + + +@withCUDA +@pytest.mark.parametrize("use_sdpa", [False, True]) +def test_multi_head_attention_bfloat16( + device: torch.device, + use_sdpa: bool, +) -> None: + attention = MultiHeadAttention( + num_heads=2, + in_features=8, + use_sdpa=use_sdpa, + device=device, + dtype=torch.bfloat16, + ) + inputs = torch.ones(2, 3, 8, device=device, dtype=torch.bfloat16) + + output = attention(inputs)[0] + + assert output.dtype == torch.bfloat16 + assert output.isfinite().all() + + +@withCUDA +def test_multi_head_attention_sdpa_fully_masked( + device: torch.device, +) -> None: + attention = MultiHeadAttention( + num_heads=2, + in_features=8, + device=device, + dtype=torch.float16, + use_sdpa=True, + ) + inputs = torch.ones(2, 3, 8, device=device, dtype=torch.float16) + patch_mask = torch.ones(2, 3, device=device, dtype=torch.bool) + + output = attention(inputs, patch_mask=patch_mask)[0] + + torch.testing.assert_close(output, torch.zeros_like(output)) + + +def test_transformer_config_attention_defaults() -> None: + transformer = TransformerConfig( + model_dims=1280, + hidden_dims=1280, + num_heads=16, + qk_norm="rms", + use_bias=False, + use_rope_seq=True, + use_rope_var=False, + ff_activation="relu", + ) + + assert transformer.use_memory_efficient_attention is True + assert transformer.use_sdpa is True + + +def _transformer_config( + *, + use_memory_efficient_attention: bool = True, + use_sdpa: bool = True, +) -> TransformerConfig: + return TransformerConfig( + model_dims=8, + hidden_dims=12, + num_heads=2, + qk_norm="rms", + use_bias=False, + use_rope_seq=True, + use_rope_var=False, + ff_activation="relu", + use_memory_efficient_attention=use_memory_efficient_attention, + use_sdpa=use_sdpa, + ) + + +@withCUDA +@pytest.mark.parametrize("use_sdpa", [False, True]) +@pytest.mark.parametrize( + ("use_memory_efficient_attention", "expected"), + [(True, 2.3272503), (False, 2.2684078)], +) +def test_mixing_transformer_uses_attention_scaling_config( + device: torch.device, + use_sdpa: bool, + use_memory_efficient_attention: bool, + expected: float, +) -> None: + config = replace( + _transformer_config( + use_memory_efficient_attention=use_memory_efficient_attention, + use_sdpa=use_sdpa, + ), + model_dims=2, + num_heads=1, + qk_norm="none", + use_rope_seq=False, + ) + transformer = MixingTransformer( + config, use_variate_attention=False, device=device + ) + with torch.no_grad(): + transformer.pre_seq_attn_ln.weight.fill_(1 / math.sqrt(2)) + for projection in ( + transformer.seq_attn.query_proj, + transformer.seq_attn.key_proj, + transformer.seq_attn.value_proj, + transformer.seq_attn.out_proj, + ): + projection.weight.copy_(torch.eye(2, device=device)) + transformer.ff1.weight.zero_() + inputs = torch.eye(2, device=device)[None, None] + patch_mask = torch.zeros(1, 1, 2, dtype=torch.bool, device=device) + + output = transformer(inputs, patch_mask)[0] + + # Full-layer values use sigmoid(1) or sigmoid(1 / sqrt(2)). + torch.testing.assert_close(output[0, 0, 1, 1], inputs.new_tensor(expected)) + + +@withCUDA +def test_mixing_transformer_residuals(device: torch.device) -> None: + transformer = MixingTransformer( + _transformer_config(), + use_variate_attention=False, + device=device, + ) + with torch.no_grad(): + transformer.seq_attn.out_proj.weight.zero_() + transformer.ff1.weight.zero_() + inputs = torch.arange( + 2 * 3 * 4 * 8, + device=device, + dtype=torch.float32, + ).reshape(2, 3, 4, 8) + patch_mask = torch.zeros(2, 3, 4, device=device, dtype=torch.bool) + + output = transformer(inputs, patch_mask)[0] + + torch.testing.assert_close(output, inputs) + + +@withCUDA +def test_mixing_transformer_routes_temporal_masks( + device: torch.device, +) -> None: + transformer = MixingTransformer( + _transformer_config(), + use_variate_attention=False, + device=device, + ) + inputs = torch.arange(48, device=device, dtype=torch.float32).reshape( + 1, 2, 3, 8 + ) + patch_mask = torch.tensor( + [[[False, True, False], [False, False, True]]], device=device + ) + + output, attention_mask = transformer(inputs, patch_mask) + torch.testing.assert_close( + attention_mask, make_attn_mask(patch_mask.reshape(2, 3)) + ) + + masked_input = inputs.clone() + masked_input[0, 0, 1, 0] += 100 + masked_output = transformer(masked_input, patch_mask)[0] + torch.testing.assert_close(masked_output[0, 0, 2], output[0, 0, 2]) + + future_input = inputs.clone() + future_input[0, 0, 2, 0] += 100 + future_output = transformer(future_input, patch_mask)[0] + torch.testing.assert_close(future_output[0, 0, 0], output[0, 0, 0]) + + +@withCUDA +def test_mixing_transformer_uses_variate_attention_and_rope( + device: torch.device, +) -> None: + rope_config = replace(_transformer_config(), use_rope_var=True) + with_rope = MixingTransformer(rope_config, device=device) + without_rope = MixingTransformer(_transformer_config(), device=device) + without_rope.load_state_dict(with_rope.state_dict()) + + with torch.no_grad(): + for transformer in (with_rope, without_rope): + transformer.seq_attn.out_proj.weight.zero_() + transformer.ff1.weight.zero_() + transformer.var_attn.query_proj.weight.copy_( + torch.eye(8, device=device) + ) + transformer.var_attn.key_proj.weight.copy_( + torch.eye(8, device=device) + ) + transformer.var_attn.value_proj.weight.copy_( + torch.eye(8, device=device) + ) + transformer.var_attn.out_proj.weight.copy_( + torch.eye(8, device=device) + ) + + inputs = torch.arange( + 3 * 2 * 8, + device=device, + dtype=torch.float32, + ).reshape(1, 3, 2, 8) + patch_mask = torch.zeros(1, 3, 2, device=device, dtype=torch.bool) + + rope_output = with_rope(inputs, patch_mask)[0] + no_rope_output = without_rope(inputs, patch_mask)[0] + + assert not torch.allclose(rope_output, inputs) + assert not torch.allclose(rope_output, no_rope_output) + + +@withCUDA +def test_mixing_transformer_variate_attention_isolates_patches_and_masks( + device: torch.device, +) -> None: + transformer = MixingTransformer(_transformer_config(), device=device) + with torch.no_grad(): + transformer.seq_attn.out_proj.weight.zero_() + transformer.ff1.weight.zero_() + for projection in ( + transformer.var_attn.query_proj, + transformer.var_attn.key_proj, + transformer.var_attn.value_proj, + transformer.var_attn.out_proj, + ): + projection.weight.copy_(torch.eye(8, device=device)) + + inputs = torch.zeros(1, 3, 2, 8, device=device) + inputs[0, 0, 0, 0] = 1 + inputs[0, 2, 0, 1] = 1 + inputs[0, 0, 1, 2] = 1 + inputs[0, 2, 1, 3] = 1 + patch_mask = torch.tensor( + [[[False, False], [False, True], [True, False]]], device=device + ) + output = transformer(inputs, patch_mask)[0] + + changed = inputs.clone() + changed[0, 0, 0, 0] = 0 + changed[0, 0, 0, 1] = 1 + changed_output = transformer(changed, patch_mask)[0] + assert not torch.allclose(changed_output[0, 1, 0], output[0, 1, 0]) + + other_patch = inputs.clone() + other_patch[0, 0, 1, 2] = 0 + other_patch[0, 0, 1, 5] = 1 + other_patch_output = transformer(other_patch, patch_mask)[0] + torch.testing.assert_close(other_patch_output[0, 1, 0], output[0, 1, 0]) + + masked = inputs.clone() + masked[0, 2, 0, 1] = 0 + masked[0, 2, 0, 4] = 1 + masked_output = transformer(masked, patch_mask)[0] + torch.testing.assert_close(masked_output[0, 1, 0], output[0, 1, 0]) + + +@withCUDA +def test_stacked_transformer_chains_layers_and_masks( + device: torch.device, +) -> None: + transformer = StackedMixingTransformer( + StackedTransformersConfig( + num_layers=2, + transformer=_transformer_config(), + ), + device=device, + ) + inputs = torch.arange(192, device=device, dtype=torch.float32).reshape( + 2, 3, 4, 8 + ) + patch_mask = torch.zeros(2, 3, 4, dtype=torch.bool, device=device) + patch_mask[0, 0, 1] = True + patch_mask[1, 2, 2] = True + + output, masks = transformer(inputs, patch_mask) + first_output, _ = transformer.layers[0](inputs, patch_mask) + expected_output, _ = transformer.layers[1](first_output, patch_mask) + expected_mask = make_attn_mask(patch_mask.reshape(2 * 3, 4)) + + torch.testing.assert_close(output, expected_output) + assert output.shape == inputs.shape + assert len(masks) == 2 + for mask in masks: + torch.testing.assert_close(mask, expected_mask) + + +@withCUDA +def test_stacked_transformer_loads_from_meta( + device: torch.device, +) -> None: + config = StackedTransformersConfig( + num_layers=2, + transformer=_transformer_config(), + ) + expected = StackedMixingTransformer(config, device=device) + transformer = StackedMixingTransformer(config, device="meta") + assert all( + buffer.device.type == "meta" for buffer in transformer.buffers() + ) + assert all( + parameter.device.type == "meta" + for parameter in transformer.parameters() + ) + + transformer.load_state_dict(expected.state_dict(), assign=True) + inputs = torch.arange( + 64, + device=device, + dtype=torch.float32, + ).reshape(1, 1, 8, 8) + patch_mask = torch.zeros(1, 1, 8, dtype=torch.bool, device=device) + + output = transformer(inputs, patch_mask)[0] + + assert all(buffer.device == device for buffer in transformer.buffers()) + assert all( + parameter.device == device for parameter in transformer.parameters() + ) + torch.testing.assert_close(output, expected(inputs, patch_mask)[0]) + + +@withCUDA +def test_multi_head_attention_manual_fully_masked_uses_finite_bias( + device: torch.device, +) -> None: + attention = MultiHeadAttention( + num_heads=2, + in_features=8, + use_rotary_position_embeddings=False, + qk_norm="none", + use_sdpa=False, + device=device, + ) + with torch.no_grad(): + for projection in ( + attention.query_proj, + attention.key_proj, + attention.value_proj, + attention.out_proj, + ): + projection.weight.copy_(torch.eye(8, device=device)) + inputs = torch.ones(1, 3, 8, device=device) + patch_mask = torch.ones(1, 3, device=device, dtype=torch.bool) + + output = attention(inputs, patch_mask=patch_mask)[0] + + torch.testing.assert_close(output, inputs) + + +@withCUDA +@pytest.mark.parametrize("use_sdpa", [False, True]) +@pytest.mark.parametrize("rescale_logits", [False, True]) +def test_multi_head_attention_matches_google_reference( + device: torch.device, + use_sdpa: bool, + rescale_logits: bool, +) -> None: + attention = MultiHeadAttention( + num_heads=2, + in_features=4, + use_sdpa=use_sdpa, + rescale_logits=rescale_logits, + device=device, + ) + with torch.no_grad(): + for projection in ( + attention.query_proj, + attention.key_proj, + attention.value_proj, + attention.out_proj, + ): + projection.weight.copy_(torch.eye(4, device=device)) + inputs = ( + torch.arange(1, 13, device=device, dtype=torch.float32).reshape( + 1, 3, 4 + ) + / 10 + ) + patch_mask = torch.tensor([[False, True, False]], device=device) + + output = attention(inputs, patch_mask=patch_mask)[0] + + # Generated with google-research/timesfm at e31dadd84cb26bd5. + expected_last = { + False: [0.828331590, 0.928331673, 1.047185063, 1.147185087], + True: [0.769981503, 0.869981468, 0.993489444, 1.093489409], + }[rescale_logits] + expected = inputs.clone() + expected[:, 1] = inputs[:, 0] + expected[:, 2] = output.new_tensor(expected_last) + torch.testing.assert_close(output, expected) + + +@withCUDA +def test_multi_head_attention_loads_from_meta(device: torch.device) -> None: + expected = MultiHeadAttention(num_heads=2, in_features=8, device=device) + attention = MultiHeadAttention( + num_heads=2, + in_features=8, + device="meta", + ) + + attention.load_state_dict(expected.state_dict(), assign=True) + inputs = torch.arange( + 24, + device=device, + dtype=torch.float32, + ).reshape(1, 3, 8) + output = attention(inputs)[0] + + assert all(buffer.device == device for buffer in attention.buffers()) + torch.testing.assert_close(output, expected(inputs)[0]) + + +@withCUDA +def test_multi_head_attention_uses_output_projection( + device: torch.device, +) -> None: + attention = MultiHeadAttention( + num_heads=2, + in_features=8, + device=device, + ) + with torch.no_grad(): + attention.out_proj.weight.zero_() + inputs = torch.arange( + 24, + device=device, + dtype=torch.float32, + ).reshape(1, 3, 8) + + output = attention(inputs)[0] + + torch.testing.assert_close(output, torch.zeros_like(output)) diff --git a/test/models/timesfm3/test_util.py b/test/models/timesfm3/test_util.py index 348ff51ec..afdfe4567 100644 --- a/test/models/timesfm3/test_util.py +++ b/test/models/timesfm3/test_util.py @@ -7,6 +7,7 @@ import torch from sdm.models.timesfm3.util import ( + get_activation_fn, get_output_patch_via_roll, get_running_stats, revin, @@ -303,6 +304,18 @@ def test_get_output_patch_via_roll_varied_sizes( assert torch.equal(wrap_mask, expected_mask[None, None]) +def test_get_activation_fn() -> None: + x = torch.tensor([-1.0, 0.0, 1.0]) + expected_silu = x * x.sigmoid() + + torch.testing.assert_close( + get_activation_fn("relu")(x), + torch.tensor([0.0, 0.0, 1.0]), + ) + torch.testing.assert_close(get_activation_fn("swish")(x), expected_silu) + assert get_activation_fn("none")(x) is x + + @withCUDA def test_stitch_patches(device: torch.device) -> None: patch_preds = torch.tensor(