diff --git a/README.md b/README.md index f949d8b..cef9d07 100644 --- a/README.md +++ b/README.md @@ -59,8 +59,8 @@ pip install -e . ```python import numpy as np import timesfm -model = timesfm.TimesFM_2p5_200M_torch() -model.load_checkpoint() +model = TimesFM_2p5_200M_torch.from_pretrained("google/timesfm-2.5-200m-pytorch") + model.compile( timesfm.ForecastConfig( max_context=1024, diff --git a/pyproject.toml b/pyproject.toml index 6055583..734a5b9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,6 +20,10 @@ huggingface_hub = { version = ">=0.23.0", extras = ["cli"] } safetensors = ">=0.5.3" torch = { version = ">=2.0.0", extras = ["cuda"] } +[tool.ruff] +line-length = 88 +indent-width = 2 + [build-system] requires = ["poetry-core"] diff --git a/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py b/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py index dfb21aa..9d201d5 100644 --- a/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py +++ b/src/timesfm/timesfm_2p5/timesfm_2p5_torch.py @@ -11,24 +11,22 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - """TimesFM models.""" import logging import math import os -from typing import Sequence +from pathlib import Path +from typing import Dict, Optional, Sequence, Union -import huggingface_hub import numpy as np -from safetensors.torch import load_file import torch +from huggingface_hub import ModelHubMixin, hf_hub_download +from safetensors.torch import load_file from torch import nn from .. import configs -from ..torch import dense -from ..torch import transformer -from ..torch import util +from ..torch import dense, transformer, util from . import timesfm_2p5_base revin = util.revin @@ -56,15 +54,17 @@ class TimesFM_2p5_200M_torch_module(nn.Module): # Layers. self.tokenizer = dense.ResidualBlock(self.config.tokenizer) - self.stacked_xf = nn.ModuleList([ + self.stacked_xf = nn.ModuleList( + [ transformer.Transformer(self.config.stacked_transformers.transformer) for _ in range(self.x) - ]) + ] + ) self.output_projection_point = dense.ResidualBlock( - self.config.output_projection_point + self.config.output_projection_point ) self.output_projection_quantiles = dense.ResidualBlock( - self.config.output_projection_quantiles + self.config.output_projection_quantiles ) # Device. @@ -82,10 +82,10 @@ class TimesFM_2p5_200M_torch_module(nn.Module): self.to(self.device) def forward( - self, - inputs: torch.Tensor, - masks: torch.Tensor, - decode_caches: list[util.DecodeCache] | None = None, + self, + inputs: torch.Tensor, + masks: torch.Tensor, + decode_caches: list[util.DecodeCache] | None = None, ): tokenizer_inputs = torch.cat([inputs, masks.to(inputs.dtype)], dim=-1) input_embeddings = self.tokenizer(tokenizer_inputs) @@ -97,17 +97,17 @@ class TimesFM_2p5_200M_torch_module(nn.Module): new_decode_caches = [] for i, layer in enumerate(self.stacked_xf): output_embeddings, new_cache = layer( - output_embeddings, masks[..., -1], decode_caches[i] + output_embeddings, masks[..., -1], decode_caches[i] ) new_decode_caches.append(new_cache) output_ts = self.output_projection_point(output_embeddings) output_quantile_spread = self.output_projection_quantiles(output_embeddings) return ( - input_embeddings, - output_embeddings, - output_ts, - output_quantile_spread, + input_embeddings, + output_embeddings, + output_ts, + output_quantile_spread, ), new_decode_caches def decode(self, horizon: int, inputs, masks): @@ -134,7 +134,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module): patch_sigma = [] for i in range(num_input_patches): (n, mu, sigma), _ = util.update_running_stats( - n, mu, sigma, patched_inputs[:, i], patched_masks[:, i] + n, mu, sigma, patched_inputs[:, i], patched_masks[:, i] ) patch_mu.append(mu) patch_sigma.append(sigma) @@ -143,47 +143,39 @@ class TimesFM_2p5_200M_torch_module(nn.Module): context_sigma = torch.stack(patch_sigma, dim=1) decode_caches = [ - util.DecodeCache( - next_index=torch.zeros( - batch_size, dtype=torch.int32, device=inputs.device - ), - num_masked=torch.zeros( - batch_size, dtype=torch.int32, device=inputs.device - ), - key=torch.zeros( - batch_size, - decode_cache_size, - self.h, - self.hd, - device=inputs.device, - ), - value=torch.zeros( - batch_size, - decode_cache_size, - self.h, - self.hd, - device=inputs.device, - ), - ) - for _ in range(self.x) + util.DecodeCache( + next_index=torch.zeros(batch_size, dtype=torch.int32, device=inputs.device), + num_masked=torch.zeros(batch_size, dtype=torch.int32, device=inputs.device), + key=torch.zeros( + batch_size, + decode_cache_size, + self.h, + self.hd, + device=inputs.device, + ), + value=torch.zeros( + batch_size, + decode_cache_size, + self.h, + self.hd, + device=inputs.device, + ), + ) + for _ in range(self.x) ] - normed_inputs = revin( - patched_inputs, context_mu, context_sigma, reverse=False - ) + normed_inputs = revin(patched_inputs, context_mu, context_sigma, reverse=False) normed_inputs = torch.where(patched_masks, 0.0, normed_inputs) (_, _, normed_outputs, normed_quantile_spread), decode_caches = self( - normed_inputs, patched_masks, decode_caches + normed_inputs, patched_masks, decode_caches ) renormed_outputs = torch.reshape( - revin(normed_outputs, context_mu, context_sigma, reverse=True), - (batch_size, -1, self.o, self.q), + revin(normed_outputs, context_mu, context_sigma, reverse=True), + (batch_size, -1, self.o, self.q), ) renormed_quantile_spread = torch.reshape( - revin( - normed_quantile_spread, context_mu, context_sigma, reverse=True - ), - (batch_size, -1, self.os, self.q), + revin(normed_quantile_spread, context_mu, context_sigma, reverse=True), + (batch_size, -1, self.os, self.q), )[:, -1, ...] # Autogressive decode @@ -192,7 +184,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module): for _ in range(num_decode_steps): new_patched_input = torch.reshape( - last_renormed_output, (batch_size, self.m, self.p) + last_renormed_output, (batch_size, self.m, self.p) ) new_mask = torch.zeros_like(new_patched_input, dtype=torch.bool) @@ -200,7 +192,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module): new_mus, new_sigmas = [], [] for i in range(self.m): (n, mu, sigma), _ = util.update_running_stats( - n, mu, sigma, new_patched_input[:, i], new_mask[:, i] + n, mu, sigma, new_patched_input[:, i], new_mask[:, i] ) new_mus.append(mu) new_sigmas.append(sigma) @@ -208,16 +200,14 @@ class TimesFM_2p5_200M_torch_module(nn.Module): new_mu = torch.stack(new_mus, dim=1) new_sigma = torch.stack(new_sigmas, dim=1) - new_normed_input = revin( - new_patched_input, new_mu, new_sigma, reverse=False - ) + new_normed_input = revin(new_patched_input, new_mu, new_sigma, reverse=False) (_, _, new_normed_output, _), decode_caches = self( - new_normed_input, new_mask, decode_caches + new_normed_input, new_mask, decode_caches ) new_renormed_output = torch.reshape( - revin(new_normed_output, new_mu, new_sigma, reverse=True), - (batch_size, self.m, self.o, self.q), + revin(new_normed_output, new_mu, new_sigma, reverse=True), + (batch_size, self.m, self.o, self.q), ) ar_outputs.append(new_renormed_output[:, -1, ...]) last_renormed_output = new_renormed_output[:, -1, :, self.aridx] @@ -230,7 +220,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module): return renormed_outputs, renormed_quantile_spread, ar_renormed_outputs def forecast_naive( - self, horizon: int, inputs: Sequence[np.ndarray] + self, horizon: int, inputs: Sequence[np.ndarray] ) -> list[np.ndarray]: """Forecasts the time series. @@ -252,11 +242,9 @@ class TimesFM_2p5_200M_torch_module(nn.Module): len_front_mask = self.p - (len(each_input) % self.p) if len_front_mask < self.p: input_t = torch.cat( - [torch.zeros(len_front_mask, dtype=torch.float32), input_t], dim=0 - ) - mask = torch.cat( - [torch.ones(len_front_mask, dtype=torch.bool), mask], dim=0 + [torch.zeros(len_front_mask, dtype=torch.float32), input_t], dim=0 ) + mask = torch.cat([torch.ones(len_front_mask, dtype=torch.bool), mask], dim=0) input_t = input_t[None, ...] mask = mask[None, ...] t_pf, _, t_ar = self.decode(horizon, input_t, mask) @@ -269,38 +257,58 @@ class TimesFM_2p5_200M_torch_module(nn.Module): return outputs -class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): +class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin): """PyTorch implementation of TimesFM 2.5 with 200M parameters.""" model: nn.Module = TimesFM_2p5_200M_torch_module() - def load_checkpoint( - self, - *, - path: str | None = None, - hf_repo_id: str | None = "google/timesfm-2.5-200m-pytorch", - ) -> None: - """Loads a PyTorch safetensors TimesFM model. - - Args: - path: Path to a local checkpoint. If not provided, will try to download - from the default Hugging Face repo. - hf_repo_id: If provided, will download from the specified Hugging Face - repo instead. + @classmethod + def _from_pretrained( + cls, + *, + model_id: str, + revision: Optional[str], + cache_dir: Optional[Union[str, Path]], + force_download: bool, + proxies: Optional[Dict], + resume_download: Optional[bool], + local_files_only: bool, + token: Optional[str], + **model_kwargs, + ): """ - if path: - pass - elif hf_repo_id: - logging.info( - "Downloading checkpoint from Hugging Face repo %s", hf_repo_id - ) - path = os.path.join( - huggingface_hub.snapshot_download(hf_repo_id), "model.safetensors" - ) - logging.info("Loading checkpoint from: %s", path) + Loads a PyTorch safetensors TimesFM model from a local path or the Hugging + Face Hub. This method is the backend for the `from_pretrained` class + method provided by `ModelHubMixin`. + """ + # Create an instance of the model wrapper class. + instance = cls(**model_kwargs) + + # Determine the path to the model weights. + model_file_path = "" + if os.path.isdir(model_id): + logging.info("Loading checkpoint from local directory: %s", model_id) + model_file_path = os.path.join(model_id, "model.safetensors") + if not os.path.exists(model_file_path): + raise FileNotFoundError(f"model.safetensors not found in directory {model_id}") else: - raise ValueError("Either path or hf_repo_id must be provided.") - self.model.load_checkpoint(path) + logging.info("Downloading checkpoint from Hugging Face repo %s", model_id) + model_file_path = hf_hub_download( + repo_id=model_id, + filename="model.safetensors", + revision=revision, + cache_dir=cache_dir, + force_download=force_download, + proxies=proxies, + resume_download=resume_download, + token=token, + local_files_only=local_files_only, + ) + + logging.info("Loading checkpoint from: %s", model_file_path) + # Load the weights into the model. + instance.model.load_checkpoint(model_file_path) + return instance def compile(self, forecast_config: configs.ForecastConfig, **kwargs) -> None: """Attempts to compile the model for fast decoding. @@ -315,7 +323,7 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): if kwargs.get("backend", None) is not None: self.model.compile(**kwargs) self.global_batch_size = ( - forecast_config.per_core_batch_size * self.model.device_count + forecast_config.per_core_batch_size * self.model.device_count ) # Shortcut. @@ -323,40 +331,36 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): if fc.max_context % self.model.p != 0: logging.info( - "When compiling, max context needs to be multiple of the patch size" - " %d. Using max context = %d instead.", - self.model.p, - new_context := math.ceil(fc.max_context / self.model.p) - * self.model.p, + "When compiling, max context needs to be multiple of the patch size" + " %d. Using max context = %d instead.", + self.model.p, + new_context := math.ceil(fc.max_context / self.model.p) * self.model.p, ) fc.max_context = new_context if fc.max_horizon % self.model.o != 0: logging.info( - "When compiling, max horizon needs to be multiple of the output patch" - " size %d. Using max horizon = %d instead.", - self.model.o, - new_horizon := math.ceil(fc.max_horizon / self.model.o) - * self.model.o, + "When compiling, max horizon needs to be multiple of the output patch" + " size %d. Using max horizon = %d instead.", + self.model.o, + new_horizon := math.ceil(fc.max_horizon / self.model.o) * self.model.o, ) fc.max_horizon = new_horizon if fc.max_context + fc.max_horizon > self.model.config.context_limit: raise ValueError( - "Context + horizon must be less than the context limit." - f" {fc.max_context} + {fc.max_horizon} >" - f" {self.model.config.context_limit}." + "Context + horizon must be less than the context limit." + f" {fc.max_context} + {fc.max_horizon} >" + f" {self.model.config.context_limit}." ) if fc.use_continuous_quantile_head and (fc.max_horizon > self.model.os): raise ValueError( - "Continuous quantile head is not supported for horizons >" - f" {self.model.os}." + f"Continuous quantile head is not supported for horizons > {self.model.os}." ) self.forecast_config = fc def _compiled_decode(horizon, inputs, masks): if horizon > fc.max_horizon: raise ValueError( - "Horizon must be less than the max horizon." - f" {horizon} > {fc.max_horizon}." + f"Horizon must be less than the max horizon. {horizon} > {fc.max_horizon}." ) inputs = torch.Tensor(np.array(inputs)).to(self.model.device) @@ -376,28 +380,25 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): mu, sigma = None, None pf_outputs, quantile_spreads, ar_outputs = self.model.decode( - forecast_config.max_horizon, inputs, masks + forecast_config.max_horizon, inputs, masks ) to_cat = [pf_outputs[:, -1, ...]] if ar_outputs is not None: to_cat.append(ar_outputs.reshape(batch_size, -1, self.model.q)) full_forecast = torch.cat(to_cat, dim=1) - flip_quantile_fn = lambda x: torch.cat( - [x[..., :1], torch.flip(x[..., 1:], dims=(-1,))], dim=-1 - ) + def flip_quantile_fn(x): + return torch.cat([x[..., :1], torch.flip(x[..., 1:], dims=(-1,))], dim=-1) if fc.force_flip_invariance: flipped_pf_outputs, flipped_quantile_spreads, flipped_ar_outputs = ( - self.model.decode(forecast_config.max_horizon, -inputs, masks) + self.model.decode(forecast_config.max_horizon, -inputs, masks) ) flipped_quantile_spreads = flip_quantile_fn(flipped_quantile_spreads) flipped_pf_outputs = flip_quantile_fn(flipped_pf_outputs) to_cat = [flipped_pf_outputs[:, -1, ...]] if flipped_ar_outputs is not None: - to_cat.append( - flipped_ar_outputs.reshape(batch_size, -1, self.model.q) - ) + to_cat.append(flipped_ar_outputs.reshape(batch_size, -1, self.model.q)) flipped_full_forecast = torch.cat(to_cat, dim=1) quantile_spreads = (quantile_spreads - flipped_quantile_spreads) / 2 pf_outputs = (pf_outputs - flipped_pf_outputs) / 2 @@ -406,30 +407,30 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): if fc.use_continuous_quantile_head: for quantile_index in [1, 2, 3, 4, 6, 7, 8, 9]: full_forecast[:, :, quantile_index] = ( - quantile_spreads[:, : fc.max_horizon, quantile_index] - - quantile_spreads[:, : fc.max_horizon, 5] - + full_forecast[:, : fc.max_horizon, 5] + quantile_spreads[:, : fc.max_horizon, quantile_index] + - quantile_spreads[:, : fc.max_horizon, 5] + + full_forecast[:, : fc.max_horizon, 5] ) full_forecast = full_forecast[:, :horizon, :] if fc.return_backcast: full_backcast = pf_outputs[:, :-1, : self.model.p, :].reshape( - batch_size, -1, self.model.q + batch_size, -1, self.model.q ) full_forecast = torch.cat([full_backcast, full_forecast], dim=1) if fc.fix_quantile_crossing: for i in [4, 3, 2, 1]: full_forecast[:, :, i] = torch.where( - full_forecast[:, :, i] < full_forecast[:, :, i + 1], - full_forecast[:, :, i], - full_forecast[:, :, i + 1], + full_forecast[:, :, i] < full_forecast[:, :, i + 1], + full_forecast[:, :, i], + full_forecast[:, :, i + 1], ) for i in [6, 7, 8, 9]: full_forecast[:, :, i] = torch.where( - full_forecast[:, :, i] > full_forecast[:, :, i - 1], - full_forecast[:, :, i], - full_forecast[:, :, i - 1], + full_forecast[:, :, i] > full_forecast[:, :, i - 1], + full_forecast[:, :, i], + full_forecast[:, :, i - 1], ) if fc.normalize_inputs: @@ -437,9 +438,9 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5): if is_positive is not None: full_forecast = torch.where( - is_positive[..., None], - torch.maximum(full_forecast, torch.zeros_like(full_forecast)), - full_forecast, + is_positive[..., None], + torch.maximum(full_forecast, torch.zeros_like(full_forecast)), + full_forecast, ) full_forecast = full_forecast.detach().cpu().numpy()