Spaces:
Sleeping
Sleeping
File size: 5,563 Bytes
37f9947 6305291 37f9947 6305291 37f9947 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 | """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()
|