Spaces:
Running on Zero
Running on Zero
File size: 4,702 Bytes
7d03019 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 | from __future__ import annotations
from dataclasses import dataclass
import hashlib
import json
import os
from pathlib import Path
import time
from huggingface_hub import hf_hub_download
@dataclass(frozen=True, slots=True)
class ModelSpec:
key: str
repo: str
filename: str
category: str
size: int
sha256: str
mount_root: str
SPECS = {
"high": ModelSpec("high", "marck1391/DaSiWa-WAN2.2-I2V-14B", "DasiwaWAN22I2V14BLightspeed_snatchkissHighV11.safetensors", "diffusion_models", 14528782272, "fa4202ea621725c57b0cbb84543bd6a5548de1d85c0c5a9f18db0bcf91202a54", "/models/high"),
"low": ModelSpec("low", "FX-FeiHou/wan2.2-Remix", "NSFW/Wan2.2_Remix_NSFW_i2v_14b_low_lighting_fp8_e4m3fn_v3.0.safetensors", "diffusion_models", 14291272136, "a239063c377e530229c4e5b1fecdcbe9de1201a30126dd42297435287e713229", "/models/low"),
"text": ModelSpec("text", "NSFW-API/NSFW-Wan-UMT5-XXL", "nsfw_wan_umt5-xxl_bf16.safetensors", "text_encoders", 11361851256, "daa3cdf38e2dc7cd776548e59d7956db3a3bc95fecf95f95e6ca2722419eeb66", "/models/text"),
"text_int8": ModelSpec("text_int8", "n-Arno/NSFW-Wan-UMT5-XXL-QINT8", "models_t5_umt5-xxl-enc-quanto_int8.safetensors", "text_encoders", 6733326448, "751f895bb238e6fc63030d107cbb7ee2e814e3793b880ab4d972139ac8b2bebb", "/models/text-int8"),
"vae": ModelSpec("vae", "Comfy-Org/Wan_2.1_ComfyUI_repackaged", "split_files/vae/wan_2.1_vae.safetensors", "vae", 253815318, "2fc39d31359a4b0a64f55876d8ff7fa8d780956ae2cb13463b0223e15148976b", "/models/vae"),
}
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(16 * 1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _resolve_source(spec: ModelSpec) -> Path:
explicit = os.environ.get(f"WAN_{spec.key.upper()}_MODEL_PATH", "").strip()
if explicit:
return Path(explicit).expanduser().resolve()
localize_active = os.environ.get(
"WAN_LOCALIZE_ACTIVE_MODELS",
os.environ.get("WAN_LOCALIZE_DIFFUSION_MODELS", "1"),
).strip().lower() not in {
"0", "false", "no", "off",
}
if spec.key in {"high", "low", "text_int8", "vae"} and localize_active:
# Model volumes remain read-only fallbacks. Their lazy/FUSE views are
# much slower when ZeroGPU rereads every active tensor for O_DIRECT
# packing. A normal ephemeral Hub cache mirrors from_pretrained and
# gives materialization and packing regular local backing files.
started = time.perf_counter()
print(f"[WAN_MODEL] {json.dumps({'event': 'localize.start', 'role': spec.key}, sort_keys=True)}", flush=True)
path = Path(hf_hub_download(
repo_id=spec.repo,
filename=spec.filename,
cache_dir=os.environ.get("WAN_HUB_CACHE", "/tmp/wan-hub-cache"),
))
print(f"[WAN_MODEL] {json.dumps({'elapsed_s': round(time.perf_counter() - started, 3), 'event': 'localize.done', 'role': spec.key}, sort_keys=True)}", flush=True)
return path
mounted = Path(spec.mount_root) / spec.filename
if mounted.is_file():
return mounted.resolve()
return Path(hf_hub_download(repo_id=spec.repo, filename=spec.filename))
def prepare_models_root(*, verify_hashes: bool | None = None) -> tuple[Path, dict[str, str]]:
# Hub-mounted weights are immutable for the selected repository revision.
# Re-hashing ~40 GB on every cold ZeroGPU request stalls the UI at model
# loading and can consume most of the GPU lease before inference starts.
verify = verify_hashes if verify_hashes is not None else os.environ.get("VERIFY_MODEL_HASHES", "0") != "0"
root = Path(os.environ.get("WAN_MODELS_ROOT", "/tmp/wan_models")).resolve()
names: dict[str, str] = {}
for spec in SPECS.values():
source = _resolve_source(spec)
if not source.is_file():
raise FileNotFoundError(f"Modello {spec.key} non trovato: {source}")
if source.stat().st_size != spec.size:
raise RuntimeError(f"Dimensione errata per {spec.key}: {source.stat().st_size}, attesa {spec.size}")
if verify and _sha256(source) != spec.sha256:
raise RuntimeError(f"SHA-256 errato per {spec.key}: {source}")
category = root / spec.category
category.mkdir(parents=True, exist_ok=True)
target = category / Path(spec.filename).name
if target.is_symlink() or target.exists():
if target.resolve() != source:
target.unlink()
if not target.exists():
target.symlink_to(source)
names[spec.key] = target.name
return root, names
|