use the ModelHubMixin api

This commit is contained in:
Kashif Rasul
2026-03-06 10:53:18 +01:00
parent 8a755c9c75
commit a8bcd5d6d9
+39 -24
View File
@@ -18,7 +18,7 @@ import logging
import math import math
import os import os
from pathlib import Path from pathlib import Path
from typing import Dict, Optional, Sequence, Union from typing import Optional, Sequence, Union
import numpy as np import numpy as np
import torch import torch
@@ -263,23 +263,43 @@ class TimesFM_2p5_200M_torch_module(nn.Module):
return outputs return outputs
class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin): class TimesFM_2p5_200M_torch(
timesfm_2p5_base.TimesFM_2p5,
ModelHubMixin,
library_name="timesfm",
repo_url="https://github.com/google-research/timesfm",
paper_url="https://arxiv.org/abs/2310.10688",
docs_url="https://github.com/google-research/timesfm",
license="apache-2.0",
pipeline_tag="time-series-forecasting",
tags=["pytorch", "timeseries", "forecasting", "timesfm-2.5"],
):
"""PyTorch implementation of TimesFM 2.5 with 200M parameters.""" """PyTorch implementation of TimesFM 2.5 with 200M parameters."""
model: nn.Module = TimesFM_2p5_200M_torch_module() DEFAULT_REPO_ID = "google/timesfm-2.5-200m-pytorch"
WEIGHTS_FILENAME = "model.safetensors"
def __init__(
self,
torch_compile: bool = True,
config: Optional[dict] = None,
):
self.model = TimesFM_2p5_200M_torch_module()
self.torch_compile = torch_compile
if config is not None:
self._hub_mixin_config = config
@classmethod @classmethod
def _from_pretrained( def _from_pretrained(
cls, cls,
*, *,
model_id: str = "google/timesfm-2.5-200m-pytorch", model_id: str = DEFAULT_REPO_ID,
revision: Optional[str], revision: Optional[str],
cache_dir: Optional[Union[str, Path]], cache_dir: Optional[Union[str, Path]],
force_download: bool = True, force_download: bool = False,
proxies: Optional[Dict] = None,
resume_download: Optional[bool] = None,
local_files_only: bool, local_files_only: bool,
token: Optional[Union[str, bool]], token: Optional[Union[str, bool]],
config: Optional[dict] = None,
**model_kwargs, **model_kwargs,
): ):
""" """
@@ -287,40 +307,35 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin):
Face Hub. This method is the backend for the `from_pretrained` class Face Hub. This method is the backend for the `from_pretrained` class
method provided by `ModelHubMixin`. method provided by `ModelHubMixin`.
""" """
# Create an instance of the model wrapper class.
instance = cls(**model_kwargs)
# Download the config file for hf tracking.
_ = hf_hub_download(
repo_id="google/timesfm-2.5-200m-pytorch",
filename="config.json",
force_download=True,
)
print("Downloaded.")
# Determine the path to the model weights. # Determine the path to the model weights.
model_file_path = "" model_file_path = ""
if os.path.isdir(model_id): if os.path.isdir(model_id):
logging.info("Loading checkpoint from local directory: %s", model_id) logging.info("Loading checkpoint from local directory: %s", model_id)
model_file_path = os.path.join(model_id, "model.safetensors") model_file_path = os.path.join(model_id, cls.WEIGHTS_FILENAME)
if not os.path.exists(model_file_path): if not os.path.exists(model_file_path):
raise FileNotFoundError(f"model.safetensors not found in directory {model_id}") raise FileNotFoundError(
f"{cls.WEIGHTS_FILENAME} not found in directory {model_id}"
)
else: else:
logging.info("Downloading checkpoint from Hugging Face repo %s", model_id) logging.info("Downloading checkpoint from Hugging Face repo %s", model_id)
model_file_path = hf_hub_download( model_file_path = hf_hub_download(
repo_id=model_id, repo_id=model_id,
filename="model.safetensors", filename=cls.WEIGHTS_FILENAME,
revision=revision, revision=revision,
cache_dir=cache_dir, cache_dir=cache_dir,
force_download=force_download, force_download=force_download,
proxies=proxies,
resume_download=resume_download,
token=token, token=token,
local_files_only=local_files_only, local_files_only=local_files_only,
) )
# Create an instance of the model wrapper class.
instance = cls(config=config, **model_kwargs)
logging.info("Loading checkpoint from: %s", model_file_path) logging.info("Loading checkpoint from: %s", model_file_path)
# Load the weights into the model. # Load the weights into the model.
instance.model.load_checkpoint(model_file_path, **model_kwargs) instance.model.load_checkpoint(
model_file_path, torch_compile=instance.torch_compile
)
return instance return instance
def _save_pretrained(self, save_directory: Union[str, Path]): def _save_pretrained(self, save_directory: Union[str, Path]):
@@ -331,7 +346,7 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin):
if not os.path.exists(save_directory): if not os.path.exists(save_directory):
os.makedirs(save_directory) os.makedirs(save_directory)
weights_path = os.path.join(save_directory, "model.safetensors") weights_path = os.path.join(save_directory, self.WEIGHTS_FILENAME)
save_file(self.model.state_dict(), weights_path) save_file(self.model.state_dict(), weights_path)
def compile(self, forecast_config: configs.ForecastConfig, **kwargs) -> None: def compile(self, forecast_config: configs.ForecastConfig, **kwargs) -> None: