Spaces:
Running
Running
| #!/usr/bin/env python3 | |
| """Download SATA demo artifacts. | |
| Models are downloaded from HuggingFace. Demo assets are downloaded from a | |
| Google Drive/direct URL zip and extracted into the repository root. | |
| """ | |
| import argparse | |
| import os | |
| import shutil | |
| import sys | |
| import tempfile | |
| import urllib.parse | |
| import urllib.request | |
| import zipfile | |
| from pathlib import Path | |
| from typing import List, Optional | |
| HF_REPO_ID = "SteveZh/sata_models" | |
| DEFAULT_ASSETS_URL = os.environ.get("SATA_ASSETS_URL", "") | |
| DEFAULT_ASSETS_GDRIVE_ID = os.environ.get("SATA_ASSETS_GDRIVE_ID", "") | |
| MODEL_PATTERNS = [ | |
| "result/vae_merge/*", | |
| "result/vae_human/*", | |
| "result/rvq_human/*", | |
| "src/mdm/save/paper_vae_human_0125/*", | |
| "src/momask-preenc/checkpoints/t2m/t2m_base/*", | |
| "src/momask-preenc/checkpoints/t2m/t2m_base/model/*", | |
| "src/momask-preenc/checkpoints/t2m/t2m_rvq/*", | |
| "src/momask-preenc/checkpoints/t2m/t2m_rvq/model/*", | |
| ] | |
| REQUIRED_MODEL_FILES = [ | |
| "result/vae_merge/config.yaml", | |
| "result/vae_merge/ms_dict.pt", | |
| "result/vae_merge/last_model.pt", | |
| "result/vae_human/config.yaml", | |
| "result/vae_human/ms_dict.pt", | |
| "result/vae_human/model_390.pt", | |
| "result/rvq_human/config.yaml", | |
| "result/rvq_human/ms_dict.pt", | |
| "result/rvq_human/last_model.pt", | |
| "src/mdm/save/paper_vae_human_0125/args.json", | |
| "src/mdm/save/paper_vae_human_0125/model000100000.pt", | |
| "src/momask-preenc/checkpoints/t2m/t2m_base/opt.txt", | |
| "src/momask-preenc/checkpoints/t2m/t2m_base/model/latest.tar", | |
| "src/momask-preenc/checkpoints/t2m/t2m_rvq/opt.txt", | |
| "src/momask-preenc/checkpoints/t2m/t2m_rvq/model/latest.tar", | |
| ] | |
| REQUIRED_ASSET_FILES = [ | |
| "data/gradio_human/bvh/000327.bvh", | |
| "data/gradio_human/processed/000327.npz", | |
| "data/gradio_animo/bvh/arctic_wolf_female_fightflee.bvh", | |
| "data/gradio_skel/human/processed/guard.npz", | |
| "data/gradio_skel/animo/processed/wisent_male_attackfence.npz", | |
| "data/gradio_target_bvh/human/guard.bvh", | |
| "data/gradio_target_bvh/animo/wisent_male_attackfence.bvh", | |
| "data/gradio_z/human/000327.pt", | |
| "data/gradio_z/animo/arctic_wolf_female_fightflee.pt", | |
| "data/test/character/default.txt", | |
| ] | |
| def repo_root() -> Path: | |
| return Path(__file__).resolve().parents[1] | |
| def missing_files(root: Path, files: List[str]) -> List[str]: | |
| return [path for path in files if not (root / path).is_file()] | |
| def verify(root: Path, *, models: bool, assets: bool) -> None: | |
| missing = [] # type: List[str] | |
| if models: | |
| missing.extend(missing_files(root, REQUIRED_MODEL_FILES)) | |
| if assets: | |
| missing.extend(missing_files(root, REQUIRED_ASSET_FILES)) | |
| if missing: | |
| print("Missing required artifact files:", file=sys.stderr) | |
| for path in missing: | |
| print(f" {path}", file=sys.stderr) | |
| raise SystemExit(1) | |
| def download_models(root: Path, repo_id: str) -> None: | |
| try: | |
| from huggingface_hub import snapshot_download | |
| except ImportError as exc: | |
| raise SystemExit( | |
| "Missing dependency: huggingface_hub. Install it with:\n" | |
| " python -m pip install huggingface-hub" | |
| ) from exc | |
| print(f"Downloading model artifacts from HuggingFace: {repo_id}") | |
| snapshot_download( | |
| repo_id=repo_id, | |
| repo_type="model", | |
| local_dir=str(root), | |
| allow_patterns=MODEL_PATTERNS, | |
| ) | |
| verify(root, models=True, assets=False) | |
| print("Model artifacts are ready.") | |
| def google_drive_url(file_id: str) -> str: | |
| return f"https://drive.google.com/uc?export=download&id={file_id}" | |
| def extract_google_drive_id(value: str) -> Optional[str]: | |
| value = value.strip() | |
| if not value: | |
| return None | |
| if "/" not in value and "?" not in value: | |
| return value | |
| parsed = urllib.parse.urlparse(value) | |
| query = urllib.parse.parse_qs(parsed.query) | |
| if "id" in query and query["id"]: | |
| return query["id"][0] | |
| parts = [part for part in parsed.path.split("/") if part] | |
| if "d" in parts: | |
| idx = parts.index("d") | |
| if idx + 1 < len(parts): | |
| return parts[idx + 1] | |
| return None | |
| def resolve_assets_url(url: str, gdrive_id: str) -> str: | |
| if gdrive_id: | |
| return google_drive_url(extract_google_drive_id(gdrive_id) or gdrive_id) | |
| if not url: | |
| raise SystemExit( | |
| "No assets URL configured. Pass --assets-url, pass --assets-gdrive-id, " | |
| "or set SATA_ASSETS_URL after uploading sata_demo_assets.zip." | |
| ) | |
| parsed = urllib.parse.urlparse(url) | |
| if "drive.google.com" in parsed.netloc: | |
| file_id = extract_google_drive_id(url) | |
| if file_id: | |
| return google_drive_url(file_id) | |
| return url | |
| def download_file(url: str, output_path: Path) -> None: | |
| print(f"Downloading assets zip: {url}") | |
| request = urllib.request.Request(url, headers={"User-Agent": "sata-artifact-downloader"}) | |
| with urllib.request.urlopen(request) as response, output_path.open("wb") as handle: | |
| shutil.copyfileobj(response, handle) | |
| def download_assets(root: Path, assets_url: str, assets_gdrive_id: str) -> None: | |
| url = resolve_assets_url(assets_url, assets_gdrive_id) | |
| with tempfile.TemporaryDirectory(prefix="sata_assets_") as tmp: | |
| zip_path = Path(tmp) / "sata_demo_assets.zip" | |
| download_file(url, zip_path) | |
| if not zipfile.is_zipfile(zip_path): | |
| raise SystemExit( | |
| "Downloaded assets file is not a valid zip. Check that the Google " | |
| "Drive file is shared publicly and that the URL points directly to " | |
| "sata_demo_assets.zip." | |
| ) | |
| print(f"Extracting {zip_path.name} into {root}") | |
| with zipfile.ZipFile(zip_path) as archive: | |
| archive.extractall(root) | |
| verify(root, models=False, assets=True) | |
| print("Demo assets are ready.") | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| actions = parser.add_argument_group("actions") | |
| actions.add_argument("--all", action="store_true", help="Download both models and assets.") | |
| actions.add_argument("--models", action="store_true", help="Download model checkpoints from HuggingFace.") | |
| actions.add_argument("--assets", action="store_true", help="Download and extract demo assets.") | |
| actions.add_argument("--check", action="store_true", help="Only verify required artifact files.") | |
| parser.add_argument("--hf-repo", default=HF_REPO_ID, help=f"HuggingFace repo id. Default: {HF_REPO_ID}") | |
| parser.add_argument("--assets-url", default=DEFAULT_ASSETS_URL, help="Direct/Google Drive URL for sata_demo_assets.zip.") | |
| parser.add_argument("--assets-gdrive-id", default=DEFAULT_ASSETS_GDRIVE_ID, help="Google Drive file id for sata_demo_assets.zip.") | |
| return parser.parse_args() | |
| def main() -> None: | |
| args = parse_args() | |
| root = repo_root() | |
| selected = args.all or args.models or args.assets or args.check | |
| if not selected: | |
| args.all = True | |
| want_models = args.all or args.models | |
| want_assets = args.all or args.assets | |
| if args.check: | |
| verify(root, models=want_models or not args.assets, assets=want_assets or not args.models) | |
| print("Artifact check passed.") | |
| return | |
| if want_models: | |
| download_models(root, args.hf_repo) | |
| if want_assets: | |
| download_assets(root, args.assets_url, args.assets_gdrive_id) | |
| verify(root, models=want_models, assets=want_assets) | |
| print("All requested artifacts are ready.") | |
| if __name__ == "__main__": | |
| main() | |