Files
AI-Profile-Router/experiments/bs-roformer-vocal-separation/app.py
T

146 lines
5.5 KiB
Python

from __future__ import annotations
import asyncio
import os
import shutil
import subprocess
import tempfile
import time
import zipfile
from pathlib import Path
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.responses import FileResponse, HTMLResponse
from starlette.background import BackgroundTask
MODEL = os.getenv("MODEL_FILENAME", "model_bs_roformer_ep_317_sdr_12.9755.ckpt")
MODEL_DIR = Path(os.getenv("MODEL_DIR", "/models"))
JOB_DIR = Path(os.getenv("JOB_DIR", "/data/jobs"))
MAX_UPLOAD = int(os.getenv("MAX_UPLOAD_BYTES", str(1024 ** 3)))
ALLOWED = {".wav", ".flac", ".mp3", ".m4a", ".aac", ".ogg", ".opus", ".wma"}
SEPARATION_LOCK = asyncio.Lock()
STARTED = time.time()
MODES = {
"vocals": {"model": MODEL, "stems": ("vocals", "instrumental"), "archive": "athena-vocals-instrumental.zip", "engine": "mdxc"},
"four_stem": {"model": "htdemucs_ft.yaml", "stems": ("vocals", "drums", "bass", "other"), "archive": "athena-4-stems.zip", "engine": "demucs"},
"six_stem": {"model": "htdemucs_6s.yaml", "stems": ("vocals", "drums", "bass", "guitar", "piano", "other"), "archive": "athena-6-stems-experimental.zip", "engine": "demucs"},
}
app = FastAPI(title="Athena Stem Separator", version="2.0")
@app.get("/", response_class=HTMLResponse)
def index() -> str:
return Path("/app/index.html").read_text(encoding="utf-8")
@app.get("/health")
def health() -> dict:
available = {name: (MODEL_DIR / mode["model"]).exists() for name, mode in MODES.items()}
return {
"status": "ok" if all(available.values()) else "starting",
"models": {name: mode["model"] for name, mode in MODES.items()},
"models_ready": available,
"busy": SEPARATION_LOCK.locked(),
"uptime_seconds": round(time.time() - STARTED, 1),
}
def _cleanup(path: Path) -> None:
shutil.rmtree(path, ignore_errors=True)
def _run_separator(input_path: Path, output_dir: Path, mode: dict) -> None:
args = [
"audio-separator", str(input_path),
"--model_filename", mode["model"],
"--model_file_dir", str(MODEL_DIR),
"--output_dir", str(output_dir),
"--output_format", "FLAC",
"--sample_rate", "44100",
"--use_soundfile",
"--use_autocast",
]
if mode["engine"] == "mdxc":
args.extend(["--mdxc_segment_size", "256", "--mdxc_overlap", "8", "--mdxc_batch_size", "1"])
else:
args.extend(["--demucs_segment_size", "40", "--demucs_shifts", "2", "--demucs_overlap", "0.25"])
completed = subprocess.run(args, capture_output=True, text=True, timeout=7200)
if completed.returncode:
detail = (completed.stderr or completed.stdout or "unknown error")[-4000:]
raise RuntimeError(detail)
def _stem_name(path: Path, expected: tuple[str, ...]) -> str | None:
lower = path.stem.lower()
for stem in sorted(expected, key=len, reverse=True):
if stem in lower:
return stem
return None
@app.post("/v1/separate")
async def separate(file: UploadFile = File(...), mode: str = Form("vocals")) -> FileResponse:
selected_mode = MODES.get(mode)
if selected_mode is None:
raise HTTPException(422, f"Unbekannter Trennmodus: {mode}")
suffix = Path(file.filename or "upload.wav").suffix.lower()
if suffix not in ALLOWED:
raise HTTPException(415, "Dieses Audioformat wird nicht unterstützt.")
if SEPARATION_LOCK.locked():
raise HTTPException(409, "Eine Trennung läuft bereits.")
job = Path(tempfile.mkdtemp(prefix="separate-", dir=JOB_DIR))
input_path = job / f"input{suffix}"
output_dir = job / "output"
output_dir.mkdir()
size = 0
try:
with input_path.open("wb") as handle:
while chunk := await file.read(1024 * 1024):
size += len(chunk)
if size > MAX_UPLOAD:
raise HTTPException(413, "Datei ist größer als 1 GiB.")
handle.write(chunk)
async with SEPARATION_LOCK:
await asyncio.to_thread(_run_separator, input_path, output_dir, selected_mode)
stems = sorted(output_dir.glob("*.flac"))
expected = selected_mode["stems"]
recognized = {_stem_name(stem, expected): stem for stem in stems}
recognized.pop(None, None)
missing = [stem for stem in expected if stem not in recognized]
if missing:
found = ", ".join(stem.name for stem in stems) or "keine"
raise RuntimeError(f"Fehlende Spuren: {', '.join(missing)}; gefunden: {found}")
archive = job / selected_mode["archive"]
with zipfile.ZipFile(archive, "w", compression=zipfile.ZIP_STORED) as bundle:
for stem in expected:
bundle.write(recognized[stem], f"{stem}.flac")
return FileResponse(
archive,
media_type="application/zip",
filename=selected_mode["archive"],
background=BackgroundTask(_cleanup, job),
)
except HTTPException:
_cleanup(job)
raise
except subprocess.TimeoutExpired:
_cleanup(job)
raise HTTPException(504, "Die Trennung hat das Zeitlimit überschritten.")
except Exception as exc:
_cleanup(job)
raise HTTPException(500, f"Trennung fehlgeschlagen: {exc}")
@app.on_event("startup")
def prepare() -> None:
JOB_DIR.mkdir(parents=True, exist_ok=True)
MODEL_DIR.mkdir(parents=True, exist_ok=True)
for old in JOB_DIR.glob("separate-*"):
if old.is_dir() and time.time() - old.stat().st_mtime > 86400:
_cleanup(old)