Spaces:
Running on Zero
Running on Zero
File size: 6,620 Bytes
1a1ea81 | 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 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | #!/usr/bin/env python3
"""Put the model payload in place before the first request.
Runs once at Space startup, from the main web process, and deliberately does
not import torch: on ZeroGPU the web process has no GPU attached, and a stray
CUDA init there is fatal for the whole Space.
Four payloads, none of which are in git:
weights/best.pt + 11 MAPS checkpoints ~76 MB private HF model repo, via
tools/fetch_assets.py
VoiceCraft EnCodec .th 1.1 GB pyp1/VoiceCraft on the Hub
audiocraft SEANet sources ~20 MB git clone; voicecraft_encodec
imports audiocraft.modules.
seanet straight from the tree
Whisper ~1.6 GB prefetched into the HF cache
The Whisper prefetch is the one that is easy to skip and shouldn't be. It is
loaded by a subprocess *inside* the GPU allocation, so downloading it there
would bill 90 seconds of the visitor's daily GPU quota to a network transfer.
Fetching it here spends startup time instead, which is free.
Every step is idempotent: a restart on a warm container re-checks and returns.
This file is copied to the Space root next to app.py by space/sync_space.sh;
REPO_ROOT below assumes that layout.
"""
from __future__ import annotations
import os
import subprocess
import sys
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parent
PINS = REPO_ROOT / "tools" / "repro_pins.env"
VOICECRAFT_ROOT = Path(
os.environ.get("VOICECRAFT_ROOT", REPO_ROOT / "third_party" / "VoiceCraft")
)
DEFAULT_ASR_MODEL = "openai/whisper-large-v3-turbo"
def _pins() -> dict[str, str]:
"""Read tools/repro_pins.env so the git refs live in exactly one place."""
out: dict[str, str] = {}
if not PINS.is_file():
return out
for line in PINS.read_text(encoding="utf-8").splitlines():
line = line.strip()
if not line or line.startswith("#") or "=" not in line:
continue
key, _, value = line.partition("=")
out[key.strip()] = value.strip().strip('"').strip("'")
return out
def _log(step: str, message: str) -> None:
print(f"[bootstrap] {step}: {message}", flush=True)
def ensure_checkpoints() -> None:
"""best.pt and the MAPS torch_models, from the private assets repo."""
fetch = REPO_ROOT / "tools" / "fetch_assets.py"
check = subprocess.run(
[sys.executable, str(fetch), "--check"], capture_output=True, text=True
)
if check.returncode == 0:
_log("checkpoints", "already present")
return
_log("checkpoints", "fetching from the assets repo")
proc = subprocess.run(
[sys.executable, str(fetch)], capture_output=True, text=True
)
if proc.returncode != 0:
raise RuntimeError(
"tools/fetch_assets.py failed. If the assets repo is still private, "
"the Space needs an HF_TOKEN secret with read access.\n"
+ (proc.stdout + proc.stderr)[-2000:]
)
_log("checkpoints", "done")
def ensure_audiocraft() -> None:
"""SEANet sources under VOICECRAFT_ROOT/src/audiocraft.
voicecraft_encodec.py imports audiocraft.modules.seanet by path rather than
installing audiocraft, so a source tree at the pinned commit is the whole
requirement — the VoiceCraft repo itself is never needed.
"""
dest = VOICECRAFT_ROOT / "src" / "audiocraft"
if (dest / "audiocraft" / "modules" / "seanet.py").is_file():
_log("audiocraft", f"already present at {dest}")
return
pins = _pins()
repo = pins.get("AUDIOCRAFT_GIT_REPO", "https://github.com/facebookresearch/audiocraft.git")
ref = pins.get("AUDIOCRAFT_GIT_REF", "main")
_log("audiocraft", f"cloning {repo}@{ref[:8]}")
dest.parent.mkdir(parents=True, exist_ok=True)
subprocess.run(
["git", "clone", "--filter=blob:none", "--quiet", repo, str(dest)], check=True
)
subprocess.run(["git", "-C", str(dest), "checkout", "--quiet", ref], check=True)
if not (dest / "audiocraft" / "modules" / "seanet.py").is_file():
raise RuntimeError(f"audiocraft SEANet sources missing under {dest}")
_log("audiocraft", "done")
def ensure_encodec() -> None:
"""The 1.1 GB VoiceCraft 16 kHz EnCodec checkpoint."""
pins = _pins()
filename = pins.get("VOICECRAFT_ENCODEC_FILE", "encodec_4cb2048_giga.th")
repo = pins.get("VOICECRAFT_ENCODEC_HF_REPO", "pyp1/VoiceCraft")
dest = VOICECRAFT_ROOT / "pretrained_models" / filename
if dest.is_file() and dest.stat().st_size > 0:
_log("encodec", "already present")
return
from huggingface_hub import hf_hub_download
_log("encodec", f"downloading {filename} from {repo}")
cached = hf_hub_download(repo_id=repo, filename=filename)
dest.parent.mkdir(parents=True, exist_ok=True)
# Copy out of the cache: the loader wants a real file at this path, and a
# cache entry can be garbage-collected underneath it.
dest.write_bytes(Path(cached).read_bytes())
_log("encodec", f"done ({dest.stat().st_size / 2**30:.1f} GiB)")
def ensure_asr(model_name: str) -> None:
"""Warm the HF cache so the transcribe subprocess never downloads."""
from huggingface_hub import snapshot_download
_log("whisper", f"prefetching {model_name}")
snapshot_download(
repo_id=model_name,
ignore_patterns=["*.bin", "*.msgpack", "*.h5", "*.ot", "*.onnx"],
)
_log("whisper", "done")
def bootstrap(asr_model: str = DEFAULT_ASR_MODEL) -> list[str]:
"""Run every step, returning the failures rather than raising.
A Space that boots and explains what is missing is far easier to debug than
one that dies during import, so failures are collected and surfaced in the
UI instead of taking the app down.
"""
problems: list[str] = []
for name, step in (
("checkpoints", ensure_checkpoints),
("audiocraft", ensure_audiocraft),
("encodec", ensure_encodec),
("whisper", lambda: ensure_asr(asr_model)),
):
try:
step()
except Exception as exc: # noqa: BLE001 - reported, not swallowed
problems.append(f"{name}: {type(exc).__name__}: {exc}")
_log(name, f"FAILED {type(exc).__name__}: {exc}")
return problems
if __name__ == "__main__":
failures = bootstrap(os.environ.get("PHONEPAINT_ASR_MODEL", DEFAULT_ASR_MODEL))
for failure in failures:
print(failure, file=sys.stderr)
sys.exit(1 if failures else 0)
|