suryatmodulus
/

GPC-1 / download.py
suryatmodulus's picture harshatheg's picture
Duplicate from harshatheg/GPC-1
96a4100
Raw History Blame Contribute Delete
8.8 kB
"""Fetch and verify the complete GPC-1 release with the standard Hub client."""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import re
import sys
from pathlib import Path, PurePosixPath
DEFAULT_REPO = "harshatheg/GPC-1"
SHA256 = re.compile(r"^[0-9a-f]{64}$")
def _release_path(root: Path, value: object, *, directory: bool) -> Path:
if not isinstance(value, str) or not value:
raise ValueError("Release metadata contains an empty path")
if "\\" in value or any(part in ("", ".", "..") for part in value.split("/")):
raise ValueError(f"Invalid release path: {value}")
relative = PurePosixPath(value)
if relative.is_absolute():
raise ValueError(f"Invalid release path: {value}")
path = root.joinpath(*relative.parts)
if not path.resolve().is_relative_to(root.resolve()):
raise ValueError(f"Release path escapes package: {value}")
if directory and not path.is_dir():
raise ValueError(f"Missing release directory: {value}")
if not directory and not path.is_file():
raise ValueError(f"Missing release file: {value}")
return path
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def verify_package(root: Path) -> dict[str, str]:
"""Validate metadata, adapter hashes, and the complete listed base-file set."""
root = root.resolve()
config_file = root / "config.json"
if not config_file.is_file() or config_file.is_symlink():
raise ValueError("Missing root config.json; fetch the full model repository")
config = json.loads(config_file.read_text(encoding="utf-8"))
if not isinstance(config, dict) or not config.get("model_type"):
raise ValueError("Root config.json is not a model config")
release = config.get("gpc1_release")
if not isinstance(release, dict) or release.get("format_version") != 1:
raise ValueError("Root config.json lacks supported gpc1_release metadata")
model = _release_path(root, release.get("model_path"), directory=True)
adapter = _release_path(root, release.get("adapter_path"), directory=True)
manifest_file = _release_path(root, release.get("runtime_manifest"), directory=False)
if model == adapter or model == root or adapter == root:
raise ValueError("Release model and adapter paths must be distinct package directories")
expected = release.get("runtime_sha256")
if not isinstance(expected, str) or not SHA256.fullmatch(expected):
raise ValueError("Invalid runtime_sha256 in release metadata")
actual = _sha256_file(manifest_file)
if actual != expected:
raise ValueError("Runtime manifest checksum mismatch")
model_config_file = model / "config.json"
if not model_config_file.is_file():
raise ValueError("Missing model/config.json")
model_config = json.loads(model_config_file.read_text(encoding="utf-8"))
root_base_config = dict(config)
root_base_config.pop("gpc1_release")
if root_base_config != model_config:
raise ValueError("Root model config does not match packaged base config")
manifest = json.loads(manifest_file.read_text(encoding="utf-8"))
if not isinstance(manifest, dict) or manifest.get("served_model_id") != "gpc-1":
raise ValueError("Runtime manifest is not for GPC-1")
assets = manifest.get("assets", {})
inventory_file = _release_path(root, assets.get("base_files"), directory=False)
if _sha256_file(inventory_file) != assets.get("base_files_sha256"):
raise ValueError("Base-file inventory checksum mismatch")
inventory = json.loads(inventory_file.read_text(encoding="utf-8"))
if (inventory.get("model_id") != manifest.get("base", {}).get("model_id")
or inventory.get("revision") != manifest.get("base", {}).get("revision")):
raise ValueError("Base-file inventory identity mismatch")
files = inventory.get("files")
if not isinstance(files, dict) or not files:
raise ValueError("Base-file inventory is empty")
for name, digest in files.items():
if not isinstance(digest, str) or not SHA256.fullmatch(digest):
raise ValueError("Invalid base-file inventory digest")
path = _release_path(model, name, directory=False)
if path.stat().st_size == 0:
raise ValueError(f"Empty base file: {name}")
if name in ("config.json", "model.safetensors.index.json") and _sha256_file(path) != digest:
raise ValueError(f"Base file checksum mismatch: {name}")
index_file = model / "model.safetensors.index.json"
if "model.safetensors.index.json" not in files or not index_file.is_file():
raise ValueError("Missing model shard index")
index = json.loads(index_file.read_text(encoding="utf-8"))
shards = set(index.get("weight_map", {}).values())
if not shards or not shards.issubset(files.keys()) or any(not isinstance(s, str) or not s.endswith(".safetensors") for s in shards):
raise ValueError("Model shard index does not match base-file inventory")
adapter_manifest = manifest.get("adapter", {})
for name, key in (("adapter_config.json", "adapter_config_sha256"),
("adapter_model.safetensors", "adapter_model_sha256")):
path = _release_path(adapter, name, directory=False)
if _sha256_file(path) != adapter_manifest.get(key):
raise ValueError(f"Adapter file checksum mismatch: {name}")
return {"root": str(root), "model": str(model), "adapter": str(adapter), "runtime_sha256": actual}
def fetch(repo: str, revision: str | None, local_dir: Path | None, cache_dir: Path | None, offline: bool) -> dict[str, str]:
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise RuntimeError("Install huggingface-hub to fetch the package") from exc
try:
path = snapshot_download(
repo_id=repo,
repo_type="model",
revision=revision,
local_dir=str(local_dir) if local_dir else None,
cache_dir=str(cache_dir) if cache_dir else None,
local_files_only=offline,
token=os.environ.get("HF_TOKEN") or None,
)
except Exception:
raise RuntimeError("Hub fetch failed; check access, revision, cache, and connectivity") from None
return verify_package(Path(path))
def counts(repo: str) -> dict[str, object]:
"""Read a Hub model-info snapshot; never fetch query files for analytics."""
try:
from huggingface_hub import HfApi
except ImportError as exc:
raise RuntimeError("Install huggingface-hub to read counts") from exc
try:
info = HfApi(token=os.environ.get("HF_TOKEN") or None).model_info(
repo_id=repo, expand=["downloads", "downloadsAllTime"]
)
except Exception:
raise RuntimeError("Hub model-info request failed; check access and connectivity") from None
return {
"repo": repo,
"downloads_30d": getattr(info, "downloads", None),
"downloads_all_time": getattr(info, "downloads_all_time", None),
}
def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description=__doc__)
sub = parser.add_subparsers(dest="command", required=True)
fetch_parser = sub.add_parser("fetch", help="Download the full model package and verify its release metadata")
fetch_parser.add_argument("--repo", default=DEFAULT_REPO)
fetch_parser.add_argument("--revision", help="Pin a tag or, preferably, a commit SHA")
fetch_parser.add_argument("--local-dir", type=Path)
fetch_parser.add_argument("--cache-dir", type=Path)
fetch_parser.add_argument("--offline", action="store_true", help="Use cached files only; no Hub request")
verify_parser = sub.add_parser("verify", help="Verify a package already on disk without network access")
verify_parser.add_argument("path", type=Path)
count_parser = sub.add_parser("counts", help="Read current Hub download counts")
count_parser.add_argument("--repo", default=DEFAULT_REPO)
args = parser.parse_args(argv)
try:
if args.command == "fetch":
result = fetch(args.repo, args.revision, args.local_dir, args.cache_dir, args.offline)
elif args.command == "verify":
result = verify_package(args.path)
else:
result = counts(args.repo)
except (OSError, ValueError, RuntimeError) as exc:
parser.exit(1, f"{args.command}: {exc}\n")
print(json.dumps(result, sort_keys=True))
return 0
if __name__ == "__main__":
sys.exit(main())