Merge pull request #372 from kashif/modelhubmixin

[HF] use the ModelHubMixin api
This commit is contained in:
Rajat Sen
2026-03-11 11:04:43 -07:00
committed by GitHub
+42 -27
View File
@@ -18,11 +18,11 @@ 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
from huggingface_hub import ModelHubMixin, hf_hub_download from huggingface_hub import PyTorchModelHubMixin, hf_hub_download
from safetensors.torch import load_file, save_file from safetensors.torch import load_file, save_file
from torch import nn from torch import nn
@@ -263,75 +263,90 @@ 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,
PyTorchModelHubMixin,
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,
): ):
""" """
Loads a PyTorch safetensors TimesFM model from a local path or the Hugging 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 Face Hub. This method is the backend for the `from_pretrained` class
method provided by `ModelHubMixin`. method provided by `PyTorchModelHubMixin`.
""" """
# 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]):
""" """
Saves the model's state dictionary to a safetensors file. This method Saves the model's state dictionary to a safetensors file. This method
is called by the `save_pretrained` method from `ModelHubMixin`. is called by the `save_pretrained` method from `PyTorchModelHubMixin`.
""" """
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: