diff --git a/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py b/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py index 5178d83..7753850 100644 --- a/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py +++ b/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py @@ -75,11 +75,19 @@ class TimesFM_2p5_200M_torch_module(nn.Module): self.device = torch.device("cpu") self.device_count = 1 - def load_checkpoint(self, path: str): + def load_checkpoint(self, path: str, **kwargs): """Loads a PyTorch TimesFM model from a checkpoint.""" tensors = load_file(path) self.load_state_dict(tensors, strict=True) self.to(self.device) + torch_compile = True + if "torch_compile" in kwargs: + torch_compile = kwargs["torch_compile"] + if torch_compile: + print("Compiling model...") + self = torch.compile(self) + + self.eval() def forward( self, @@ -310,7 +318,7 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin): logging.info("Loading checkpoint from: %s", model_file_path) # Load the weights into the model. - instance.model.load_checkpoint(model_file_path) + instance.model.load_checkpoint(model_file_path, **model_kwargs) return instance def _save_pretrained(self, save_directory: Union[str, Path]): @@ -333,10 +341,6 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin): forecast_config: Configuration for forecasting flags. **kwargs: Additional keyword arguments to pass to model.compile(). """ - - if forecast_config.torch_compile: - self.model = torch.compile(self.model) - self.model.eval() self.global_batch_size = ( forecast_config.per_core_batch_size * self.model.device_count ) diff --git a/src/timesfm/torch/transformer.py b/src/timesfm/torch/transformer.py index 3b14a3c..b7fba8b 100644 --- a/src/timesfm/torch/transformer.py +++ b/src/timesfm/torch/transformer.py @@ -129,6 +129,28 @@ def _dot_product_attention( return torch.einsum("...hqk,...khd->...qhd", attn_weights, value) +def _torch_dot_product_attention(query, key, value, mask=None): + """ + Performs the exact same (unscaled) attention as your original function, + but using the fast and fused F.scaled_dot_product_attention kernel. + """ + + # 1. Permute inputs from (B, L, H, D) to the expected (B, H, L, D) + query = query.permute(0, 2, 1, 3) + key = key.permute(0, 2, 1, 3) + value = value.permute(0, 2, 1, 3) + + # 2. Call the fused attention kernel + # - Pass the mask to `attn_mask`. + # - Set `scale=1.0` to disable the default 1/sqrt(d_k) scaling. + output = F.scaled_dot_product_attention(query, key, value, attn_mask=mask, scale=1.0) + + # 3. Permute the output back to the original (B, L, H, D) layout + output = output.permute(0, 2, 1, 3) + + return output + + class PerDimScale(nn.Module): """Per-dimension scaling.""" @@ -155,7 +177,7 @@ class MultiHeadAttention(nn.Module): use_per_dim_scale: bool = True, use_rotary_position_embeddings: bool = True, use_bias: bool = False, - attention_fn: Callable[..., torch.Tensor] = _dot_product_attention, + attention_fn: Callable[..., torch.Tensor] = _torch_dot_product_attention, qk_norm: str = "rms", fuse_qkv: bool = False, ):