Team Ai
Apppublic

commitcopilot/infer-003

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
download_model.py78 linesDownload Raw Back to scripts
1import argparse2import json3import shutil4import urllib.request5from pathlib import Path6from typing import Any7 8from huggingface_hub import hf_hub_download9 10GGUF_MAGIC = b"GGUF"11 12 13def load_config(path: Path) -> dict[str, Any]:14    with path.open("r", encoding="utf-8") as file:15        return json.load(file)16 17 18def download_from_hub(config: dict[str, Any], output_path: Path, token: str | None) -> None:19    repo_id = str(config["repo_id"])20    filename = str(config["filename"])21    revision = str(config.get("revision", "main"))22    downloaded = hf_hub_download(23        repo_id=repo_id,24        filename=filename,25        revision=revision,26        token=token,27    )28    shutil.copyfile(downloaded, output_path)29 30 31def download_from_url(config: dict[str, Any], output_path: Path) -> None:32    url = str(config["url"])33    with urllib.request.urlopen(url) as response:34        with output_path.open("wb") as file:35            shutil.copyfileobj(response, file)36 37 38def validate_gguf(path: Path) -> None:39    if not path.exists() or path.stat().st_size == 0:40        raise RuntimeError(f"Downloaded model is empty: {path}")41 42    with path.open("rb") as file:43        magic = file.read(4)44    if magic != GGUF_MAGIC:45        raise RuntimeError(f"Downloaded file is not a GGUF model: {path}")46 47 48def main() -> None:49    parser = argparse.ArgumentParser()50    parser.add_argument("--config", default="config.json")51    parser.add_argument("--output-dir", default="/models")52    parser.add_argument("--hf-token", default="")53    args = parser.parse_args()54 55    config = load_config(Path(args.config))56    download_config = config.get("download")57    if not isinstance(download_config, dict):58        raise RuntimeError("config.download must be an object")59 60    output_dir = Path(args.output_dir)61    output_dir.mkdir(parents=True, exist_ok=True)62    output_filename = str(download_config.get("output_filename", "model.gguf"))63    output_path = output_dir / output_filename64 65    source = str(download_config.get("source", "hf_hub"))66    if source == "hf_hub":67        token = args.hf_token.strip() or None68        download_from_hub(download_config, output_path, token)69    elif source == "url":70        download_from_url(download_config, output_path)71    else:72        raise RuntimeError(f"Unsupported download source: {source}")73 74    validate_gguf(output_path)75 76 77if __name__ == "__main__":78    main()