commitcopilot/infer-003
0
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()