Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
236 changes: 236 additions & 0 deletions sdm/models/timesfm3/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
57 changes: 57 additions & 0 deletions test/models/timesfm3/reference/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]:
Expand Down Expand Up @@ -114,13 +141,43 @@ 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,
"forward": {
"inputs": FORWARD_INPUTS,
"outputs": forward_outputs,
},
"decode_without_stitching": {
"config_overrides": no_stitching_overrides,
"inputs": NO_STITCHING_INPUTS,
"output": no_stitching.tolist(),
},
}


Expand Down
Loading
Loading