HandEdit-LoRA / scripts /download_weights.py
HandEdit's picture
Add files using upload-large-folder tool
ce47bc4 verified
Raw History Blame
2.63 kB
#!/usr/bin/env python3
"""Download and safely extract the published HandEdit LoRA archive."""
from __future__ import annotations
import argparse
import hashlib
import json
import shutil
import zipfile
from pathlib import Path
from huggingface_hub import hf_hub_download
DEFAULT_REPO_ID = "HandEdit/HandEdit-LoRA"
ARCHIVE_NAME = "checkpoints.zip"
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(8 * 1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def safe_extract(archive: Path, destination: Path) -> None:
root = destination.resolve()
with zipfile.ZipFile(archive) as bundle:
for member in bundle.infolist():
target = (root / member.filename).resolve()
if target != root and root not in target.parents:
raise RuntimeError(f"Unsafe archive member: {member.filename}")
bundle.extractall(root)
def main() -> None:
parser = argparse.ArgumentParser(
description="Download the sanitized HandEdit LoRA checkpoints."
)
parser.add_argument("--repo-id", default=DEFAULT_REPO_ID)
parser.add_argument(
"--output-dir",
type=Path,
default=Path("."),
help="Repository root where checkpoints/ will be extracted.",
)
parser.add_argument(
"--keep-archive",
action="store_true",
help="Keep checkpoints.zip after successful extraction.",
)
args = parser.parse_args()
output_dir = args.output_dir.expanduser().resolve()
output_dir.mkdir(parents=True, exist_ok=True)
archive = Path(
hf_hub_download(
repo_id=args.repo_id,
filename=ARCHIVE_NAME,
repo_type="model",
local_dir=output_dir,
)
)
manifest_path = Path(
hf_hub_download(
repo_id=args.repo_id,
filename="weights_manifest.json",
repo_type="model",
local_dir=output_dir,
)
)
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
expected = manifest["archive"]["sha256"]
actual = sha256(archive)
if actual != expected:
raise RuntimeError(
f"Archive checksum mismatch: expected {expected}, received {actual}"
)
safe_extract(archive, output_dir)
print(f"[OK] Extracted sanitized weights to: {output_dir / 'checkpoints'}")
if not args.keep_archive:
archive.unlink()
print(f"[OK] Removed downloaded archive: {archive}")
if __name__ == "__main__":
main()