use snapshot_download

This commit is contained in:
Kashif Rasul
2024-05-10 11:32:53 +02:00
parent afeb84f77a
commit d5f55b0825
+2 -2
View File
@@ -24,7 +24,7 @@ import jax
import jax.numpy as jnp import jax.numpy as jnp
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from huggingface_hub import hf_hub_download from huggingface_hub import snapshot_download
from paxml import checkpoints from paxml import checkpoints
from paxml import tasks_lib from paxml import tasks_lib
from praxis import base_hyperparams from praxis import base_hyperparams
@@ -235,7 +235,7 @@ class TimesFm:
step: step of the checkpoint to load. If `None`, load lastest checkpoint. step: step of the checkpoint to load. If `None`, load lastest checkpoint.
""" """
# Download the checkpoint from Hugging Face Hub # Download the checkpoint from Hugging Face Hub
checkpoint_path = hf_hub_download(repo_id) checkpoint_path = snapshot_download(repo_id)
# Initialize the model weights. # Initialize the model weights.
self._logging("Constructing model weights.") self._logging("Constructing model weights.")