use the ModelHubMixin api
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user