Files
2026-07-13 12:47:19 +08:00

137 lines
5.9 KiB
Python

# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file.
import importlib.util
import os
from contextlib import contextmanager
from pathlib import Path
from litgpt.config import configs
from litgpt.constants import _HF_TRANSFER_AVAILABLE, _SAFETENSORS_AVAILABLE
from litgpt.scripts.convert_hf_checkpoint import convert_hf_checkpoint
def download_from_hub(
repo_id: str,
access_token: str | None = os.getenv("HF_TOKEN"),
tokenizer_only: bool = False,
convert_checkpoint: bool = True,
dtype: str | None = None,
checkpoint_dir: Path = Path("checkpoints"),
model_name: str | None = None,
) -> None:
"""Download weights or tokenizer data from the Hugging Face Hub.
Arguments:
repo_id: The repository ID in the format ``org/name`` or ``user/name`` as shown in Hugging Face.
If "list" is provided as input, a list of the currently supported models in LitGPT and quits.
access_token: Optional API token to access models with restrictions.
tokenizer_only: Whether to download only the tokenizer files.
convert_checkpoint: Whether to convert the checkpoint files to the LitGPT format after downloading.
dtype: The data type to convert the checkpoint files to. If not specified, the weights will remain in the
dtype they are downloaded in.
checkpoint_dir: Where to save the downloaded files.
model_name: The existing config name to use for this repo_id. This is useful to download alternative weights of
existing architectures.
"""
options = [f"{config['hf_config']['org']}/{config['hf_config']['name']}" for config in configs]
if repo_id == "list":
print("Please specify --repo_id <repo_id>. Available values:")
print("\n".join(sorted(options, key=lambda x: x.lower())))
return
if model_name is None and repo_id not in options:
print(
f"Unsupported `repo_id`: {repo_id}."
"\nIf you are trying to download alternative "
"weights for a supported model, please specify the corresponding model via the `--model_name` option, "
"for example, `litgpt download NousResearch/Hermes-2-Pro-Llama-3-8B --model_name Llama-3-8B`."
"\nAlternatively, please choose a valid `repo_id` from the list of supported models, which can be obtained via "
"`litgpt download list`."
)
return
from huggingface_hub import snapshot_download
if importlib.util.find_spec("hf_transfer") is None:
print(
"It is recommended to install hf_transfer for faster checkpoint download speeds: `pip install hf_transfer`"
)
download_files = ["tokenizer*", "generation_config.json", "config.json"]
if not tokenizer_only:
bins, safetensors = find_weight_files(repo_id, access_token)
if bins:
# covers `.bin` files and `.bin.index.json`
download_files.append("*.bin*")
elif safetensors:
if not _SAFETENSORS_AVAILABLE:
raise ModuleNotFoundError(str(_SAFETENSORS_AVAILABLE))
download_files.append("*.safetensors*")
else:
raise ValueError(f"Couldn't find weight files for {repo_id}")
import huggingface_hub._snapshot_download as download
import huggingface_hub.constants as constants
previous = constants.HF_HUB_ENABLE_HF_TRANSFER
if _HF_TRANSFER_AVAILABLE and not previous:
print("Setting HF_HUB_ENABLE_HF_TRANSFER=1")
constants.HF_HUB_ENABLE_HF_TRANSFER = True
download.HF_HUB_ENABLE_HF_TRANSFER = True
directory = checkpoint_dir / repo_id
with gated_repo_catcher(repo_id, access_token):
snapshot_download(
repo_id,
local_dir=directory,
allow_patterns=download_files,
token=access_token,
)
constants.HF_HUB_ENABLE_HF_TRANSFER = previous
download.HF_HUB_ENABLE_HF_TRANSFER = previous
if convert_checkpoint and not tokenizer_only:
print("Converting checkpoint files to LitGPT format.")
convert_hf_checkpoint(checkpoint_dir=directory, dtype=dtype, model_name=model_name)
def find_weight_files(repo_id: str, access_token: str | None) -> tuple[list[str], list[str]]:
from huggingface_hub import repo_info
from huggingface_hub.utils import filter_repo_objects
with gated_repo_catcher(repo_id, access_token):
info = repo_info(repo_id, token=access_token)
filenames = [f.rfilename for f in info.siblings]
bins = list(filter_repo_objects(items=filenames, allow_patterns=["*model*.bin*"]))
safetensors = list(filter_repo_objects(items=filenames, allow_patterns=["*.safetensors*"]))
return bins, safetensors
@contextmanager
def gated_repo_catcher(repo_id: str, access_token: str | None):
try:
yield
except OSError as e:
err_msg = str(e)
if "Repository Not Found" in err_msg:
raise ValueError(
f"Repository at https://huggingface.co/api/models/{repo_id} not found."
" Please make sure you specified the correct `repo_id`."
) from None
elif "gated repo" in err_msg:
if not access_token:
raise ValueError(
f"https://huggingface.co/{repo_id} requires authentication, please set the `HF_TOKEN=your_token`"
" environment variable or pass `--access_token=your_token`. You can find your token by visiting"
" https://huggingface.co/settings/tokens."
) from None
else:
raise ValueError(
f"https://huggingface.co/{repo_id} requires authentication. The access token provided by `HF_TOKEN=your_token`"
" environment variable or `--access_token=your_token` may not have sufficient access rights. Please"
f" visit https://huggingface.co/{repo_id} for more information."
) from None
raise e from None