feat: add multi-stem audio separation modes

This commit is contained in:
Mikei386 committed 2026-09-09 00:30:08 +02:00
1 parent 30203bf13b
commit e2f35517f8
6 files changed
+76 -40

No files matched your search

+42 -21
View File
@@ -9,7 +9,7 @@ import time
import zipfile
from pathlib import Path
from fastapi import FastAPI, File, HTTPException, UploadFile
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.responses import FileResponse, HTMLResponse
from starlette.background import BackgroundTask
@@ -22,7 +22,13 @@ 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")
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)
@@ -32,11 +38,11 @@ def index() -> str:
@app.get("/health")
def health() -> dict:
checkpoint = MODEL_DIR / MODEL
available = {name: (MODEL_DIR / mode["model"]).exists() for name, mode in MODES.items()}
return {
"status": "ok" if checkpoint.exists() else "starting",
"model": MODEL,
"model_ready": checkpoint.exists(),
"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),
}
@@ -46,28 +52,40 @@ def _cleanup(path: Path) -> None:
shutil.rmtree(path, ignore_errors=True)
def _run_separator(input_path: Path, output_dir: Path) -> None:
def _run_separator(input_path: Path, output_dir: Path, mode: dict) -> None:
args = [
"audio-separator", str(input_path),
"--model_filename", MODEL,
"--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",
"--mdxc_segment_size", "256",
"--mdxc_overlap", "8",
"--mdxc_batch_size", "1",
]
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(...)) -> FileResponse:
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.")
@@ -87,21 +105,24 @@ async def separate(file: UploadFile = File(...)) -> FileResponse:
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)
await asyncio.to_thread(_run_separator, input_path, output_dir, selected_mode)
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"
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 stems:
lower = stem.name.lower()
target = "vocals.flac" if "vocal" in lower else "instrumental.flac"
bundle.write(stem, target)
for stem in expected:
bundle.write(recognized[stem], f"{stem}.flac")
return FileResponse(
archive,
media_type="application/zip",
filename="athena-vocals-instrumental.zip",
filename=selected_mode["archive"],
background=BackgroundTask(_cleanup, job),
)
except HTTPException: