diff --git a/sdm/models/timesfm3/core.py b/sdm/models/timesfm3/core.py index 76761d20a..eb5afac91 100644 --- a/sdm/models/timesfm3/core.py +++ b/sdm/models/timesfm3/core.py @@ -29,6 +29,9 @@ 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 ( @@ -254,3 +257,104 @@ def _preprocess( (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/test/models/timesfm3/test_core.py b/test/models/timesfm3/test_core.py index ec973c659..9f9878e31 100644 --- a/test/models/timesfm3/test_core.py +++ b/test/models/timesfm3/test_core.py @@ -1,6 +1,7 @@ # 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 ( @@ -9,6 +10,7 @@ TransformerConfig, ) from sdm.models.timesfm3.core import _TimesFM3Model +from sdm.models.timesfm3.util import get_running_stats from sdm.testing import withCUDA @@ -167,3 +169,270 @@ def test_internal_model_preserves_covariates_under_cpm( [[[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:])