From eb9786e3390c38ce0432b59ec7c9af3b9f7b4670 Mon Sep 17 00:00:00 2001 From: Filippo Guerranti Date: Wed, 30 Sep 2026 07:35:54 +0000 Subject: [PATCH] feat(timesfm3): add baseline full-sequence decoding --- sdm/models/timesfm3/core.py | 236 +++++++++++++++++++++ test/models/timesfm3/reference/generate.py | 57 +++++ test/models/timesfm3/reference/golden.json | 162 +++++++++++++- test/models/timesfm3/test_core.py | 88 ++++++++ 4 files changed, 537 insertions(+), 6 deletions(-) diff --git a/sdm/models/timesfm3/core.py b/sdm/models/timesfm3/core.py index eb5afac91..d5b760546 100644 --- a/sdm/models/timesfm3/core.py +++ b/sdm/models/timesfm3/core.py @@ -358,3 +358,239 @@ def forward( outputs["__call__:seq_attn_mask"] = attention_masks outputs["__call__:transformer_output"] = transformer_output return outputs + + def decode( + self, + target: Tensor, + horizon: int = 0, + past_only_covariates: Tensor | None = None, + past_future_covariates: Tensor | None = None, + target_mask: Tensor | None = None, + past_only_mask: Tensor | None = None, + past_future_mask: Tensor | None = None, + mask: Tensor | None = None, + return_aux_outputs: bool = False, + ) -> Tensor | tuple[Tensor, dict[str, Any]]: + """Decode a forecast in one non-autoregressive pass. + + Future-known covariates determine the forecast horizon when supplied, + overriding ``horizon``. + + Gradient tracking follows the ambient PyTorch gradient mode. + + Args: + target: Target context with shape ``[B, U, C]``, where ``B`` is the + batch size and ``U`` is the number of target variates. + ``C`` is the context length. + horizon: Number of future time steps to predict. + past_only_covariates: Historical covariates with shape + ``[B, V, C]``. + past_future_covariates: Future-known covariates with shape + ``[B, W, C + H]``, where ``H`` is the forecast horizon. + target_mask: Invalid-value mask with shape ``[B, U, C]``. + past_only_mask: Invalid-value mask with shape ``[B, V, C]``. + past_future_mask: Invalid-value mask with shape + ``[B, W, C + H]``. + mask: Global context mask with shape ``[B, C]``. + return_aux_outputs: Whether to return full-sequence intermediate + outputs with the forecast. + + Returns: + Forecasts for every target and covariate variate with shape + ``[B, U + V + W, H, Q]``, where ``Q`` is the number of quantiles, + optionally paired with the full-sequence outputs. + """ + device = target.device + batch_size, num_target, context = target.shape + + if past_future_covariates is not None: + horizon = past_future_covariates.shape[-1] - context + if horizon <= 0: + raise ValueError("Decode function requires horizon > 0.") + if self.use_stitching or self.use_linear_detrending: + raise NotImplementedError( + "Stitched decoding and linear detrending are not yet " + "implemented. Set both options to False." + ) + + # 1. Pad context to multiple of input_patch_len + ctx_padding = ( + self.input_patch_len - (context % self.input_patch_len) + ) % self.input_patch_len + if ctx_padding > 0: + target = torch.nn.functional.pad(target, (ctx_padding, 0)) + if mask is not None: + mask = torch.nn.functional.pad( + mask, (ctx_padding, 0), value=True + ) + if past_only_covariates is not None: + past_only_covariates = torch.nn.functional.pad( + past_only_covariates, (ctx_padding, 0) + ) + if past_future_covariates is not None: + past_future_covariates = torch.nn.functional.pad( + past_future_covariates, (ctx_padding, 0) + ) + if target_mask is not None: + target_mask = torch.nn.functional.pad( + target_mask, (ctx_padding, 0), value=True + ) + if past_only_mask is not None: + past_only_mask = torch.nn.functional.pad( + past_only_mask, (ctx_padding, 0), value=True + ) + if past_future_mask is not None: + past_future_mask = torch.nn.functional.pad( + past_future_mask, (ctx_padding, 0), value=True + ) + context = context + ctx_padding + + if mask is None: + mask = torch.zeros( + batch_size, context, dtype=torch.bool, device=device + ) + if ctx_padding > 0: + mask[:, :ctx_padding] = True + + # 2. Pad horizon + hor_padding = (-horizon) % self.output_patch_len + padded_horizon = horizon + hor_padding + num_horizon_patches = padded_horizon // self.input_patch_len + num_context_patches = context // self.input_patch_len + + # 3. Build context & horizon inputs + if target_mask is None: + target_mask = torch.zeros_like(target, dtype=torch.bool) + target_mask = target_mask | mask.unsqueeze(1) + + all_ctx_vals = [target] + all_ctx_masks = [target_mask] + num_past_only = 0 + if past_only_covariates is not None: + num_past_only = past_only_covariates.shape[1] + if past_only_mask is None: + past_only_mask = torch.zeros_like( + past_only_covariates, dtype=torch.bool + ) + all_ctx_vals.append(past_only_covariates) + all_ctx_masks.append(past_only_mask | mask.unsqueeze(1)) + if past_future_covariates is not None: + if past_future_mask is None: + past_future_mask = torch.zeros_like( + past_future_covariates, dtype=torch.bool + ) + all_ctx_vals.append(past_future_covariates[..., :context]) + all_ctx_masks.append( + past_future_mask[..., :context] | mask.unsqueeze(1) + ) + + ctx_vals = torch.cat(all_ctx_vals, dim=1) + ctx_masks = torch.cat(all_ctx_masks, dim=1) + + ctx_vals = torch.where(ctx_masks, 0.0, ctx_vals) + + all_hor_vals = [ + torch.zeros(batch_size, num_target, padded_horizon, device=device), + torch.zeros( + batch_size, num_past_only, padded_horizon, device=device + ), + ] + all_hor_masks = [ + torch.ones( + batch_size, + num_target, + padded_horizon, + dtype=torch.bool, + device=device, + ), + torch.ones( + batch_size, + num_past_only, + padded_horizon, + dtype=torch.bool, + device=device, + ), + ] + + if past_future_covariates is not None: + if past_future_mask is None: + past_future_mask = torch.zeros_like( + past_future_covariates, dtype=torch.bool + ) + pf_future_vals = past_future_covariates[ + ..., context : context + horizon + ] + pf_future_masks = past_future_mask[ + ..., context : context + horizon + ] + pf_future_vals = torch.where(pf_future_masks, 0.0, pf_future_vals) + if hor_padding > 0: + pf_future_vals = torch.nn.functional.pad( + pf_future_vals, (0, hor_padding) + ) + pf_future_masks = torch.nn.functional.pad( + pf_future_masks, (0, hor_padding), value=True + ) + all_hor_vals.append(pf_future_vals) + all_hor_masks.append(pf_future_masks) + + hor_vals = torch.cat(all_hor_vals, dim=1) + hor_masks = torch.cat(all_hor_masks, dim=1) + + all_vals = torch.cat([ctx_vals, hor_vals], dim=-1) + all_masks = torch.cat([ctx_masks, hor_masks], dim=-1) + + num_variates = all_vals.shape[1] + patch_is_target = torch.zeros( + ( + batch_size, + num_variates, + num_context_patches + num_horizon_patches, + ), + dtype=torch.bool, + device=device, + ) + patch_is_target[:, : num_target + num_past_only, :] = True + + # Reshape values & masks to patched shape (b, v, n, p) + values_bvnp = all_vals.reshape( + batch_size, num_variates, -1, self.input_patch_len + ) + masks_bvnp = all_masks.reshape( + batch_size, num_variates, -1, self.input_patch_len + ) + + # Build horizon CPM mask: context=False, horizon=True. + num_total_patches = num_context_patches + num_horizon_patches + horizon_cpm_mask = torch.zeros( + batch_size, num_total_patches, dtype=torch.bool, device=device + ) + horizon_cpm_mask[:, num_context_patches:] = True + + freeze_after = ( + num_context_patches - 1 if self.use_frozen_running_stats else None + ) + forward_out = self.forward( + values_bvnp, + masks_bvnp, + patch_is_target, + freeze_after=freeze_after, + patch_cpm_mask=horizon_cpm_mask, + return_aux_outputs=return_aux_outputs, + ) + logits = forward_out[ + "logits" + ] # (b, v, n, output_patch_len, num_quantiles) + + num_forecast_chunks = padded_horizon // self.output_patch_len + forecast_indices = torch.arange( + num_forecast_chunks, device=device + ) * self.rolls + (num_context_patches - 1) + forecast_logits = logits[:, :, forecast_indices, :, :] + horizon_logits = forecast_logits.reshape( + batch_size, num_variates, -1, self.num_quantiles + )[:, :, :horizon, :] + + if return_aux_outputs: + return horizon_logits, forward_out + return horizon_logits diff --git a/test/models/timesfm3/reference/generate.py b/test/models/timesfm3/reference/generate.py index 47c224593..0157c259f 100644 --- a/test/models/timesfm3/reference/generate.py +++ b/test/models/timesfm3/reference/generate.py @@ -59,6 +59,33 @@ "patch_is_target": [[[True, True, True]]], "patch_cpm_masks": [None, [[False, True, True]]], } +NO_STITCHING_INPUTS: dict[str, Any] = { + "target": [[[1.0, 2.0, 5.0, 3.0, 8.0]]], + "past_only_covariates": [[[3.0, 1.0, 4.0, 1.0, 5.0]]], + "past_future_covariates": [ + [[5.0, 6.0, 4.0, 9.0, 7.0, 8.0, 6.0, 10.0, 12.0, 7.0, 15.0]] + ], + "mask": [[False, False, False, True, False]], + "past_only_mask": [[[False, True, False, False, False]]], + "past_future_mask": [ + [ + [ + False, + False, + False, + False, + False, + False, + True, + False, + False, + True, + False, + ] + ] + ], + "horizon": 99, +} def _fill_state(model: torch.nn.Module) -> list[dict[str, Any]]: @@ -114,6 +141,31 @@ def _internal_fixture(upstream: Any) -> dict[str, Any]: )["logits"] forward_outputs.append(logits.tolist()) + no_stitching_overrides = { + "use_stitching": False, + "use_linear_detrending": False, + "use_frozen_running_stats": True, + } + no_stitching_config = {**INTERNAL_CONFIG, **no_stitching_overrides} + no_stitching_model = upstream.TimesFM3Torch(**no_stitching_config).eval() + _fill_state(no_stitching_model) + with torch.inference_mode(): + no_stitching = no_stitching_model.decode( + torch.tensor(NO_STITCHING_INPUTS["target"]), + horizon=NO_STITCHING_INPUTS["horizon"], + past_only_covariates=torch.tensor( + NO_STITCHING_INPUTS["past_only_covariates"] + ), + past_future_covariates=torch.tensor( + NO_STITCHING_INPUTS["past_future_covariates"] + ), + mask=torch.tensor(NO_STITCHING_INPUTS["mask"]), + past_only_mask=torch.tensor(NO_STITCHING_INPUTS["past_only_mask"]), + past_future_mask=torch.tensor( + NO_STITCHING_INPUTS["past_future_mask"] + ), + ) + return { "config": INTERNAL_CONFIG, "state": state, @@ -121,6 +173,11 @@ def _internal_fixture(upstream: Any) -> dict[str, Any]: "inputs": FORWARD_INPUTS, "outputs": forward_outputs, }, + "decode_without_stitching": { + "config_overrides": no_stitching_overrides, + "inputs": NO_STITCHING_INPUTS, + "output": no_stitching.tolist(), + }, } diff --git a/test/models/timesfm3/reference/golden.json b/test/models/timesfm3/reference/golden.json index aca0a5cab..be7f2e432 100644 --- a/test/models/timesfm3/reference/golden.json +++ b/test/models/timesfm3/reference/golden.json @@ -169,16 +169,16 @@ [ [ [ - 1, - 3 + 1.0, + 3.0 ], [ - 0, - 0 + 0.0, + 0.0 ], [ - 0, - 0 + 0.0, + 0.0 ] ] ] @@ -319,6 +319,156 @@ ] ] ] + }, + "decode_without_stitching": { + "config_overrides": { + "use_stitching": false, + "use_linear_detrending": false, + "use_frozen_running_stats": true + }, + "inputs": { + "target": [ + [ + [ + 1.0, + 2.0, + 5.0, + 3.0, + 8.0 + ] + ] + ], + "past_only_covariates": [ + [ + [ + 3.0, + 1.0, + 4.0, + 1.0, + 5.0 + ] + ] + ], + "past_future_covariates": [ + [ + [ + 5.0, + 6.0, + 4.0, + 9.0, + 7.0, + 8.0, + 6.0, + 10.0, + 12.0, + 7.0, + 15.0 + ] + ] + ], + "mask": [ + [ + false, + false, + false, + true, + false + ] + ], + "past_only_mask": [ + [ + [ + false, + true, + false, + false, + false + ] + ] + ], + "past_future_mask": [ + [ + [ + false, + false, + false, + false, + false, + false, + true, + false, + false, + true, + false + ] + ] + ], + "horizon": 99 + }, + "output": [ + [ + [ + [ + 2.878720283508301 + ], + [ + 3.1826422214508057 + ], + [ + 3.4865641593933105 + ], + [ + 3.7904860973358154 + ], + [ + 3.0529794692993164 + ], + [ + 3.1188242435455322 + ] + ], + [ + [ + 3.672028064727783 + ], + [ + 3.755284070968628 + ], + [ + 3.8385400772094727 + ], + [ + 3.9217963218688965 + ], + [ + 3.713366985321045 + ], + [ + 3.7319252490997314 + ] + ], + [ + [ + 4.564489364624023 + ], + [ + 4.984493255615234 + ], + [ + 5.404496669769287 + ], + [ + 5.824500560760498 + ], + [ + 4.63043737411499 + ], + [ + 5.022459983825684 + ] + ] + ] + ] } } } diff --git a/test/models/timesfm3/test_core.py b/test/models/timesfm3/test_core.py index 33551a91f..313918f41 100644 --- a/test/models/timesfm3/test_core.py +++ b/test/models/timesfm3/test_core.py @@ -49,6 +49,7 @@ def _internal_model( *, use_iterative_cpm_revin: bool = True, use_linear_detrending: bool = True, + use_stitching: bool = True, ) -> _TimesFM3Model: return _TimesFM3Model( input_patch_len=2, @@ -58,6 +59,7 @@ def _internal_model( transformer_config=_transformer_config(), use_iterative_cpm_revin=use_iterative_cpm_revin, use_linear_detrending=use_linear_detrending, + use_stitching=use_stitching, device=device, dtype=dtype, ) @@ -526,3 +528,89 @@ def test_internal_model_matches_pinned_upstream( expected = actual.new_tensor(fixture["outputs"][case_index]) torch.testing.assert_close(actual, expected) + + +@withCUDA +def test_internal_model_decode_without_stitching_matches_pinned_upstream( + device: torch.device, +) -> None: + reference = _reference_fixture() + fixture = reference["internal"]["decode_without_stitching"] + inputs = fixture["inputs"] + config = { + **reference["internal"]["config"], + **fixture["config_overrides"], + } + model = _TimesFM3Model(**_model_config(config), device=device).eval() + _load_reference_weights( + model, reference["internal"], reference["weight_recipe"] + ) + result = model.decode( + torch.tensor(inputs["target"], device=device), + horizon=inputs["horizon"], + past_only_covariates=torch.tensor( + inputs["past_only_covariates"], device=device + ), + past_future_covariates=torch.tensor( + inputs["past_future_covariates"], device=device + ), + past_only_mask=torch.tensor(inputs["past_only_mask"], device=device), + past_future_mask=torch.tensor( + inputs["past_future_mask"], device=device + ), + mask=torch.tensor(inputs["mask"], device=device), + return_aux_outputs=True, + ) + + assert isinstance(result, tuple) + actual, auxiliary = result + expected = actual.new_tensor(fixture["output"]) + torch.testing.assert_close(actual, expected) + assert actual.shape == (1, 3, 6, 1) # Covariates override horizon=99. + assert auxiliary["logits"].shape == (1, 3, 7, 4, 1) + assert auxiliary["__call__:resblock_input"].shape == (1, 3, 7, 12) + + +@withCUDA +def test_internal_model_decode_pads_target_only_horizon( + device: torch.device, +) -> None: + model = _internal_model( + device, use_stitching=False, use_linear_detrending=False + ).eval() + forecast = model.decode(torch.ones(1, 1, 4, device=device), horizon=5) + + assert isinstance(forecast, torch.Tensor) + assert forecast.shape == (1, 1, 5, 3) + assert torch.isfinite(forecast).all() + + +def test_internal_model_decode_rejects_nonpositive_horizon() -> None: + model = _golden_model(torch.device("cpu")) + + with pytest.raises(ValueError, match="horizon > 0"): + model.decode(torch.zeros(1, 1, 2)) + + +@withCUDA +def test_internal_decode_supports_autograd(device: torch.device) -> None: + model = _internal_model( + device, + use_linear_detrending=False, + use_stitching=False, + ).train() + target = torch.arange( + 1, + 9, + dtype=torch.float32, + device=device, + ).reshape(1, 1, 8) + target.requires_grad_() + + forecast = model.decode(target, horizon=3) + assert isinstance(forecast, torch.Tensor) + forecast.sum().backward() + + assert target.grad is not None + assert torch.isfinite(target.grad).all() + assert any(parameter.grad is not None for parameter in model.parameters())