feat: add target-and-remainder stem extraction
This commit is contained in:
1 parent
0f0e77928a
commit
2342495d24
4 files changed
+85
-26
No files matched your search
@@ -27,6 +27,13 @@ MODES = {
|
||||
"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"},
|
||||
}
|
||||
TARGETS = {
|
||||
"vocals": {"mode": "vocals", "stem": "vocals", "remainder": "instrumental", "archive": "athena-gesang-und-rest.zip", "rest_file": "instrumental.flac"},
|
||||
"drums": {"mode": "four_stem", "stem": "drums", "archive": "athena-schlagzeug-und-rest.zip", "rest_file": "rest-ohne-schlagzeug.flac"},
|
||||
"bass": {"mode": "four_stem", "stem": "bass", "archive": "athena-bass-und-rest.zip", "rest_file": "rest-ohne-bass.flac"},
|
||||
"guitar": {"mode": "six_stem", "stem": "guitar", "archive": "athena-gitarre-und-rest.zip", "rest_file": "rest-ohne-gitarre.flac"},
|
||||
"piano": {"mode": "six_stem", "stem": "piano", "archive": "athena-piano-und-rest.zip", "rest_file": "rest-ohne-piano.flac"},
|
||||
}
|
||||
|
||||
app = FastAPI(title="Athena Stem Separator", version="2.0")
|
||||
|
||||
@@ -43,6 +50,7 @@ def health() -> dict:
|
||||
"status": "ok" if all(available.values()) else "starting",
|
||||
"models": {name: mode["model"] for name, mode in MODES.items()},
|
||||
"models_ready": available,
|
||||
"targets": list(TARGETS),
|
||||
"busy": SEPARATION_LOCK.locked(),
|
||||
"uptime_seconds": round(time.time() - STARTED, 1),
|
||||
}
|
||||
@@ -80,11 +88,34 @@ def _stem_name(path: Path, expected: tuple[str, ...]) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _mix_remainder(stems: list[Path], output_path: Path) -> None:
|
||||
args = ["ffmpeg", "-hide_banner", "-loglevel", "error", "-y"]
|
||||
for stem in stems:
|
||||
args.extend(["-i", str(stem)])
|
||||
inputs = "".join(f"[{index}:a]" for index in range(len(stems)))
|
||||
args.extend([
|
||||
"-filter_complex", f"{inputs}amix=inputs={len(stems)}:normalize=0:dropout_transition=0[rest]",
|
||||
"-map", "[rest]", "-ar", "44100", "-c:a", "flac", str(output_path),
|
||||
])
|
||||
completed = subprocess.run(args, capture_output=True, text=True, timeout=1800)
|
||||
if completed.returncode:
|
||||
detail = (completed.stderr or completed.stdout or "unknown ffmpeg error")[-4000:]
|
||||
raise RuntimeError(f"Restspur konnte nicht erzeugt werden: {detail}")
|
||||
|
||||
|
||||
@app.post("/v1/separate")
|
||||
async def separate(file: UploadFile = File(...), mode: str = Form("vocals")) -> FileResponse:
|
||||
selected_mode = MODES.get(mode)
|
||||
async def separate(
|
||||
file: UploadFile = File(...),
|
||||
target: str | None = Form(None),
|
||||
mode: str | None = Form(None),
|
||||
) -> FileResponse:
|
||||
selected_target = TARGETS.get(target) if target else None
|
||||
if target and selected_target is None:
|
||||
raise HTTPException(422, f"Unbekannte Zielspur: {target}")
|
||||
selected_mode_name = selected_target["mode"] if selected_target else (mode or "vocals")
|
||||
selected_mode = MODES.get(selected_mode_name)
|
||||
if selected_mode is None:
|
||||
raise HTTPException(422, f"Unbekannter Trennmodus: {mode}")
|
||||
raise HTTPException(422, f"Unbekannter Trennmodus: {selected_mode_name}")
|
||||
suffix = Path(file.filename or "upload.wav").suffix.lower()
|
||||
if suffix not in ALLOWED:
|
||||
raise HTTPException(415, "Dieses Audioformat wird nicht unterstützt.")
|
||||
@@ -114,14 +145,30 @@ async def separate(file: UploadFile = File(...), mode: str = Form("vocals")) ->
|
||||
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"]
|
||||
archive_name = selected_target["archive"] if selected_target else selected_mode["archive"]
|
||||
archive = job / archive_name
|
||||
with zipfile.ZipFile(archive, "w", compression=zipfile.ZIP_STORED) as bundle:
|
||||
for stem in expected:
|
||||
bundle.write(recognized[stem], f"{stem}.flac")
|
||||
if selected_target:
|
||||
target_stem = selected_target["stem"]
|
||||
bundle.write(recognized[target_stem], f"{target_stem}.flac")
|
||||
if "remainder" in selected_target:
|
||||
remainder = recognized[selected_target["remainder"]]
|
||||
else:
|
||||
remainder = job / selected_target["rest_file"]
|
||||
await asyncio.to_thread(
|
||||
_mix_remainder,
|
||||
[recognized[stem] for stem in expected if stem != target_stem],
|
||||
remainder,
|
||||
)
|
||||
bundle.write(remainder, selected_target["rest_file"])
|
||||
else:
|
||||
# Rückwärtskompatibilität für bestehende API-Clients.
|
||||
for stem in expected:
|
||||
bundle.write(recognized[stem], f"{stem}.flac")
|
||||
return FileResponse(
|
||||
archive,
|
||||
media_type="application/zip",
|
||||
filename=selected_mode["archive"],
|
||||
filename=archive_name,
|
||||
background=BackgroundTask(_cleanup, job),
|
||||
)
|
||||
except HTTPException:
|
||||
|
||||
Reference in new issue
Block a user