# 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 . 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