APUS-OpenJev-v1 / download_model.py
gump2049's picture
Fix family download metadata and counted download flows (#2)
b008be4
Raw History Blame
2.89 kB
"""Download one APUS-OpenJev variant using its verified family catalog.
The Hub's meta.yaml is JSON (a YAML 1.2 subset), so no YAML dependency is needed.
Only the requested model subdirectory is downloaded. Requires huggingface_hub.
"""
import argparse
import hashlib
import json
from pathlib import Path
import re
REPO_ID = "apus-ailab/APUS-OpenJev-v1"
VARIANTS = ("4B-5949", "9B-3000", "9B-5949")
SCHEMA = "apus-openjev-family.v1"
def download_variant(*, local_dir, variant="9B-3000", revision="main"):
"""Resolve once, validate the catalog, download and verify the model config."""
if variant not in VARIANTS:
raise ValueError("Unknown model variant")
from huggingface_hub import HfApi, hf_hub_download, snapshot_download
commit = HfApi().model_info(REPO_ID, revision=revision).sha
if not isinstance(commit, str) or not re.fullmatch(r"[0-9a-f]{40}", commit):
raise ValueError("Hub did not return an immutable commit SHA")
catalog_path = hf_hub_download(REPO_ID, "meta.yaml", revision=commit)
catalog = json.loads(Path(catalog_path).read_text(encoding="utf-8"))
if (
not isinstance(catalog, dict)
or catalog.get("schema") != SCHEMA
or catalog.get("repo_id") != REPO_ID
or not isinstance(catalog.get("variants"), dict)
):
raise ValueError("Invalid model family catalog")
entry = catalog["variants"].get(variant)
if not isinstance(entry, dict) or entry.get("subfolder") != variant:
raise ValueError("Catalog subfolder does not match the selected variant")
expected_sha = entry.get("config_sha256")
if not isinstance(expected_sha, str) or not re.fullmatch(
r"[0-9a-f]{64}", expected_sha
):
raise ValueError("Catalog config_sha256 is invalid")
snapshot = snapshot_download(
repo_id=REPO_ID,
revision=commit,
allow_patterns=[f"{variant}/*"],
local_dir=str(local_dir),
)
model_path = Path(snapshot) / variant
actual_sha = hashlib.sha256((model_path / "config.json").read_bytes()).hexdigest()
if actual_sha != expected_sha:
raise ValueError("Downloaded config.json does not match the family catalog")
return {
"repo_id": REPO_ID,
"revision": commit,
"variant": variant,
"model_path": str(model_path.resolve()),
"config_sha256": actual_sha,
}
def main(argv=None):
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=VARIANTS, default="9B-3000")
parser.add_argument("--revision", default="main")
parser.add_argument("--local-dir", type=Path, required=True)
args = parser.parse_args(argv)
result = download_variant(
local_dir=args.local_dir, variant=args.variant, revision=args.revision
)
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()