InstructAV2AV / space_model_store.py
suimu's picture
init
e0177dc
Raw
History Blame Contribute Delete
4.61 kB
"""Lazy Hugging Face Hub model downloads for the InstructAV2AV Space."""
from __future__ import annotations
import logging
import os
from pathlib import Path
from threading import Lock
from typing import Any
from huggingface_hub import hf_hub_download, snapshot_download
INSTRUCTAV2AV_REPO = os.getenv(
"INSTRUCTAV2AV_REPO", "suimu/InstructAV2AV"
)
WAN_REPO = os.getenv("INSTRUCTAV2AV_WAN_REPO", "Wan-AI/Wan2.2-TI2V-5B")
MMAUDIO_REPO = os.getenv("INSTRUCTAV2AV_MMAUDIO_REPO", "hkchengrex/MMAudio")
WAN_FILES = (
"models_t5_umt5-xxl-enc-bf16.pth",
"Wan2.2_VAE.pth",
)
MMAUDIO_FILES = (
"ext_weights/v1-16.pth",
"ext_weights/best_netG.pt",
)
def _hub_token() -> str | None:
return os.getenv("HF_TOKEN") or os.getenv("HUGGING_FACE_HUB_TOKEN") or None
class _NoOpProgress:
def __call__(self, *_args: Any, **_kwargs: Any) -> None:
return None
class HubModelStore:
"""Create the checkpoint layout expected by the upstream inference code."""
def __init__(
self,
model_home: str | Path,
hub_cache_dir: str | Path | None = None,
):
self.model_home = Path(model_home).expanduser().resolve()
self.hub_cache_dir = Path(
hub_cache_dir or self.model_home / "hub"
).expanduser().resolve()
self.ckpt_dir = self.model_home / "ckpts"
self.hub_cache_dir.mkdir(parents=True, exist_ok=True)
self.ckpt_dir.mkdir(parents=True, exist_ok=True)
self._base_ready = self._base_layout_ready()
self._lock = Lock()
def _base_layout_ready(self) -> bool:
wan_dir = self.ckpt_dir / "Wan2.2-TI2V-5B"
mmaudio_dir = self.ckpt_dir / "MMAudio"
return (
all((wan_dir / filename).is_file() for filename in WAN_FILES)
and (wan_dir / "google" / "umt5-xxl").is_dir()
and all((mmaudio_dir / filename).is_file() for filename in MMAUDIO_FILES)
)
@staticmethod
def _link(source: Path, target: Path) -> None:
if target.exists():
return
if target.is_symlink():
target.unlink()
if not source.exists():
raise FileNotFoundError(f"Downloaded model file is missing: {source}")
target.parent.mkdir(parents=True, exist_ok=True)
target.symlink_to(source, target_is_directory=source.is_dir())
def ensure_base_models(self, progress: Any) -> None:
if self._base_ready:
return
with self._lock:
if self._base_ready:
return
token = _hub_token()
progress(0.01, desc="Downloading Wan text encoder and video VAE…")
wan_snapshot = Path(
snapshot_download(
repo_id=WAN_REPO,
allow_patterns=[*WAN_FILES, "google/*"],
cache_dir=self.hub_cache_dir,
token=token,
)
)
wan_target = self.ckpt_dir / "Wan2.2-TI2V-5B"
for filename in WAN_FILES:
self._link(wan_snapshot / filename, wan_target / filename)
self._link(wan_snapshot / "google", wan_target / "google")
progress(0.04, desc="Downloading MMAudio VAE and vocoder…")
mmaudio_snapshot = Path(
snapshot_download(
repo_id=MMAUDIO_REPO,
allow_patterns=list(MMAUDIO_FILES),
cache_dir=self.hub_cache_dir,
token=token,
)
)
mmaudio_target = self.ckpt_dir / "MMAudio"
for filename in MMAUDIO_FILES:
self._link(mmaudio_snapshot / filename, mmaudio_target / filename)
self._base_ready = True
logging.info("Shared model files are ready under %s", self.ckpt_dir)
def resolve_checkpoint(self, model_key: str, progress: Any) -> Path:
self.ensure_base_models(progress)
filename = f"{model_key}.safetensors"
progress(0.06, desc=f"Downloading {model_key} editing checkpoint…")
checkpoint = Path(
hf_hub_download(
repo_id=INSTRUCTAV2AV_REPO,
filename=filename,
cache_dir=self.hub_cache_dir,
token=_hub_token(),
)
)
logging.info("Editing checkpoint is ready: %s", checkpoint)
return checkpoint
def preload_default(self) -> Path:
"""Synchronously cache shared weights and the default general checkpoint."""
return self.resolve_checkpoint("general", _NoOpProgress())