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
104 changes: 104 additions & 0 deletions sdm/models/timesfm3/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
Loading
Loading