#!/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)