137 lines
5.9 KiB
Python
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
|