Sever / app.py
Radiotomy's picture
Use stable demucs 4.0.1 API (get_model + apply_model)
f3f8ae9 verified
Raw History Blame Contribute Delete
5.05 kB
"""
Sever — BASE Station's own stem separation engine.
Runs HTDemucs 6-source (htdemucs_6s) locally instead of paying a third-party
API per separation. Six stems: drums, bass, other, vocals, guitar, piano.
Design notes that are load-bearing:
* SEPARATION IS ASYNCHRONOUS. A 3-minute track on CPU takes minutes, far past
any sane HTTP timeout, so /separate returns a job id immediately and the
caller polls /status. Doing it inline would guarantee gateway timeouts.
* ONE JOB AT A TIME. max_workers=1 is deliberate. Concurrent separations each
hold a full model + a full audio tensor in RAM, which is exactly how the
Coda and Siren Song Spaces earn their 'meta tensor' / OOM crashes. Queueing
is slower but never fails; parallelism here would be a memory bomb.
* MODEL LOADED ONCE, LAZILY. Loading htdemucs_6s costs a weights download and
~2GB of RAM. At import time that would fail the Space boot health check; per
request it would multiply the cost by every job.
* LOW-LEVEL demucs API ON PURPOSE. `demucs.api.Separator` exists only on the
unreleased branch — on the pinned 4.0.1 release the real entry points are
get_model + apply_model, so that is what this uses.
* soundfile, NOT torchaudio.save. torchaudio's native backend (torchcodec)
does not load on the HF Space image — an already-paid-for lesson from Skye.
"""
import os
import urllib.request
import uuid
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Dict, List, Optional
import soundfile as sf
import torch
from fastapi import FastAPI
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel
OUT = Path("/tmp/outputs")
OUT.mkdir(parents=True, exist_ok=True)
MODEL_NAME = "htdemucs_6s"
app = FastAPI(title="Sever", description="HTDemucs-6s stem separation")
# Static mount so a finished stem is a plain downloadable URL. The backend
# copies these into Base44 storage — a /tmp file on a sleeping Space is gone.
app.mount("/outputs", StaticFiles(directory=str(OUT)), name="outputs")
JOBS: Dict[str, dict] = {}
POOL = ThreadPoolExecutor(max_workers=1)
_model = None
def get_model_once():
global _model
if _model is None:
from demucs.pretrained import get_model
m = get_model(MODEL_NAME)
m.cpu()
m.eval()
_model = m
return _model
class SeparateRequest(BaseModel):
audio_url: str
stems: Optional[List[str]] = None
def _run(job_id: str, audio_url: str, want: Optional[List[str]]):
from demucs.apply import apply_model
from demucs.audio import AudioFile
src = OUT / f"{job_id}_src"
try:
JOBS[job_id]["status"] = "processing"
urllib.request.urlretrieve(audio_url, src)
model = get_model_once()
wav = AudioFile(str(src)).read(
streams=0, samplerate=model.samplerate, channels=model.audio_channels
)
# Demucs expects the input normalized to zero mean / unit variance, then
# the same transform undone on the output. Skipping this does not error —
# it just quietly degrades separation quality.
ref = wav.mean(0)
wav = (wav - ref.mean()) / ref.std()
with torch.no_grad():
sources = apply_model(model, wav[None], device="cpu", progress=False)[0]
sources = sources * ref.std() + ref.mean()
files = {}
for name, tensor in zip(model.sources, sources):
if want and name not in want:
continue
path = OUT / f"{job_id}_{name}.wav"
# (channels, samples) -> (samples, channels) for soundfile
sf.write(str(path), tensor.t().cpu().numpy(), model.samplerate)
files[name] = f"/outputs/{path.name}"
JOBS[job_id].update(
status="completed",
stems=files,
sample_rate=model.samplerate,
model=MODEL_NAME,
)
except Exception as exc: # surfaced to the caller instead of a silent stall
JOBS[job_id].update(status="failed", error=f"{type(exc).__name__}: {exc}")
finally:
if src.exists():
try:
os.remove(src)
except OSError:
pass
@app.get("/health")
def health():
return {
"ok": True,
"model": MODEL_NAME,
"stems": list(_model.sources) if _model is not None else None,
"loaded": _model is not None,
"queued": len([j for j in JOBS.values() if j["status"] == "queued"]),
}
@app.post("/separate")
def separate(req: SeparateRequest):
if not req.audio_url:
return {"error": "audio_url is required"}
job_id = uuid.uuid4().hex
JOBS[job_id] = {"status": "queued", "stems": {}}
POOL.submit(_run, job_id, req.audio_url, req.stems)
return {"job_id": job_id, "status": "queued", "model": MODEL_NAME}
@app.get("/status/{job_id}")
def status(job_id: str):
job = JOBS.get(job_id)
if not job:
return {"error": "unknown job_id", "status": "failed"}
return {"job_id": job_id, **job}