Add BS-RoFormer vocal separation mode
This commit is contained in:
1 parent
a4e894fe70
commit
0069b61dbb
15 files changed
+421
-52
No files matched your search
@@ -0,0 +1,124 @@
|
||||
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, 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()
|
||||
|
||||
app = FastAPI(title="Athena Vocal Separator", version="1.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:
|
||||
checkpoint = MODEL_DIR / MODEL
|
||||
return {
|
||||
"status": "ok" if checkpoint.exists() else "starting",
|
||||
"model": MODEL,
|
||||
"model_ready": checkpoint.exists(),
|
||||
"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) -> None:
|
||||
args = [
|
||||
"audio-separator", str(input_path),
|
||||
"--model_filename", MODEL,
|
||||
"--model_file_dir", str(MODEL_DIR),
|
||||
"--output_dir", str(output_dir),
|
||||
"--output_format", "FLAC",
|
||||
"--sample_rate", "44100",
|
||||
"--use_soundfile",
|
||||
"--use_autocast",
|
||||
"--mdxc_segment_size", "256",
|
||||
"--mdxc_overlap", "8",
|
||||
"--mdxc_batch_size", "1",
|
||||
]
|
||||
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)
|
||||
|
||||
|
||||
@app.post("/v1/separate")
|
||||
async def separate(file: UploadFile = File(...)) -> FileResponse:
|
||||
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)
|
||||
|
||||
stems = sorted(output_dir.glob("*.flac"))
|
||||
if len(stems) != 2:
|
||||
raise RuntimeError(f"Erwartet wurden zwei FLAC-Dateien, gefunden: {len(stems)}")
|
||||
archive = job / "athena-vocals-instrumental.zip"
|
||||
with zipfile.ZipFile(archive, "w", compression=zipfile.ZIP_STORED) as bundle:
|
||||
for stem in stems:
|
||||
lower = stem.name.lower()
|
||||
target = "vocals.flac" if "vocal" in lower else "instrumental.flac"
|
||||
bundle.write(stem, target)
|
||||
return FileResponse(
|
||||
archive,
|
||||
media_type="application/zip",
|
||||
filename="athena-vocals-instrumental.zip",
|
||||
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)
|
||||
Reference in new issue
Block a user