chora-backend / separator.py
k25kar's picture
perf: mp3 stems + throttled progress
6305291
Raw
History Blame Contribute Delete
5.56 kB
"""Demucs source-separation wrapper.
Splits a mixed stereo file into 4 stems (vocals / drums / bass / other)
using htdemucs. Used both as a library (by the FastAPI service) and as a
CLI for one-off separation:
python separator.py song.mp3 -o separated/
"""
from __future__ import annotations
import os
# Must be set before torch is imported: a few htdemucs ops are not yet
# implemented on Apple's MPS backend and need the CPU fallback.
os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")
import argparse
import logging
import time
from pathlib import Path
from typing import Callable
import torch
from demucs.api import Separator, save_audio
logger = logging.getLogger(__name__)
STEMS = ("vocals", "drums", "bass", "other")
DEFAULT_MODEL = "htdemucs"
# progress in [0.0, 1.0]
ProgressCallback = Callable[[float], None]
def pick_device() -> str:
"""Best available torch device: cuda > mps > cpu."""
if torch.cuda.is_available():
return "cuda"
if torch.backends.mps.is_available():
return "mps"
return "cpu"
def _demucs_progress_adapter(on_progress: ProgressCallback) -> Callable[[dict], None]:
"""Adapt demucs' raw callback dict to a single 0-1 float.
Demucs reports per-segment offsets per model in the bag; normalize
across (models x audio_length) so the caller sees monotonic progress.
"""
def callback(data: dict) -> None:
try:
models = max(int(data.get("models", 1)), 1)
model_idx = int(data.get("model_idx_in_bag", 0))
length = max(int(data.get("audio_length", 1)), 1)
offset = int(data.get("segment_offset", 0))
if data.get("state") == "end":
offset = min(offset + int(data.get("segment_length", 0) or 0), length)
fraction = (model_idx * length + offset) / (models * length)
on_progress(min(max(fraction, 0.0), 1.0))
except Exception: # progress must never kill a separation job
logger.debug("progress callback failed", exc_info=True)
return callback
def separate(
input_path: str | Path,
output_dir: str | Path,
model: str = DEFAULT_MODEL,
device: str | None = None,
shifts: int = 0,
overlap: float = 0.25,
output_format: str = "wav", # "wav" | "mp3" (mp3 is ~10x smaller, 320kbps)
on_progress: ProgressCallback | None = None,
) -> dict[str, Path]:
"""Separate ``input_path`` into 4 stem wav files under ``output_dir``.
Returns a mapping of stem name -> written file path.
Raises FileNotFoundError / RuntimeError on bad input or model failure.
"""
input_path = Path(input_path)
if not input_path.is_file():
raise FileNotFoundError(f"No such audio file: {input_path}")
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
device = device or pick_device()
logger.info("Separating %s with %s on %s", input_path.name, model, device)
separator = Separator(
model=model,
device=device,
shifts=shifts,
overlap=overlap,
callback=_demucs_progress_adapter(on_progress) if on_progress else None,
)
start = time.perf_counter()
_origin, separated = separator.separate_audio_file(input_path)
elapsed = time.perf_counter() - start
logger.info("Separation finished in %.1fs", elapsed)
written: dict[str, Path] = {}
for stem, tensor in separated.items():
out_path = output_dir / f"{stem}.{output_format}"
if output_format == "mp3":
save_audio(tensor, str(out_path), samplerate=separator.samplerate, bitrate=320)
else:
save_audio(tensor, str(out_path), samplerate=separator.samplerate)
written[stem] = out_path
logger.info("Wrote %s", out_path)
missing = set(STEMS) - set(written)
if missing:
raise RuntimeError(f"Model did not produce expected stems: {missing}")
return written
def main() -> None:
parser = argparse.ArgumentParser(description="Separate a stereo track into 4 stems with Demucs.")
parser.add_argument("input", help="Path to .mp3/.wav file")
parser.add_argument("-o", "--output", default="separated", help="Output directory (default: separated/)")
parser.add_argument("-m", "--model", default=DEFAULT_MODEL, help=f"Demucs model (default: {DEFAULT_MODEL})")
parser.add_argument("-d", "--device", default=None, help="torch device: cuda | mps | cpu (default: auto)")
parser.add_argument(
"--shifts", type=int, default=0,
help="random time-shift passes to average (reduces artifacts, each adds a full pass; default: 0)",
)
parser.add_argument(
"--overlap", type=float, default=0.25,
help="segment overlap 0-0.99 (higher = smoother seams, slower; default: 0.25)",
)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
last_shown = -1
def show_progress(fraction: float) -> None:
nonlocal last_shown
pct = int(fraction * 100)
if pct >= last_shown + 5: # print every 5%
last_shown = pct
print(f" progress: {pct}%", flush=True)
stems = separate(
args.input,
args.output,
model=args.model,
device=args.device,
shifts=args.shifts,
overlap=args.overlap,
on_progress=show_progress,
)
print("\nDone. Stems written:")
for name, path in stems.items():
print(f" {name:>7}: {path}")
if __name__ == "__main__":
main()