PhonePaint / bootstrap.py
Tom-Brenner's picture
PhonePaint Space: MAPS-only, ZeroGPU, assets fetched at startup
1a1ea81 verified
Raw History Blame Contribute Delete
6.62 kB
#!/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)