File size: 8,221 Bytes
39ea11c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
"""
Pure numpy + onnxruntime reference implementation for the HT-Demucs FT
drums specialist. NO TORCH at inference — works in any environment with
numpy, soundfile, and onnxruntime.

Usage:
    python infer.py input.mp3 out_dir/
    # writes out_dir/drums.wav (the meaningful stem)

Or as a library:
    import infer
    drums = infer.separate_drums("song.mp3", "htdemucs_ft_drums.onnx")
    # drums: numpy array (channels, samples) at 44.1 kHz

This script was hand-rolled to be readable and portable rather than maximally
fast. For production deployment you'll want to:
  - Use IOBinding to skip the numpy<->ORT copies
  - Use GraphOptimizationLevel.ORT_ENABLE_ALL
  - Add CoreMLExecutionProvider / DmlExecutionProvider / CUDAExecutionProvider
    depending on platform (see https://onnxruntime.ai/docs/execution-providers/)
"""
from __future__ import annotations

import argparse
import sys
import time
from pathlib import Path

import numpy as np
import onnxruntime as ort
import soundfile as sf

# These constants are baked into the exported graph. If you re-export with
# a different segment length, update them here.
SAMPLE_RATE = 44100
SEGMENT_S = 7.8
N_SAMPLES = int(SEGMENT_S * SAMPLE_RATE)  # 343,980
N_CHANNELS = 2
SOURCES = ["drums", "bass", "other", "vocals"]  # the bag's output order
SPECIALIST_STEM = "drums"  # which stem is meaningfully predicted by this model
DEFAULT_ONNX = Path(__file__).resolve().parent / "htdemucs_ft_drums.onnx"


def _make_transition_window(segment: int, overlap_frac: float = 0.25) -> np.ndarray:
    transition = int(segment * overlap_frac)
    window = np.ones(segment, dtype=np.float32)
    fade = np.linspace(0, 1, transition, dtype=np.float32)
    window[:transition] = fade
    window[-transition:] = fade[::-1]
    return window


def _load_session(onnx_path: Path,
                  providers: list[str] | None = None) -> ort.InferenceSession:
    if providers is None:
        # CPU-only default. To enable CoreML on macOS, pass:
        #   providers=["CoreMLExecutionProvider", "CPUExecutionProvider"]
        # (first-time CoreML compile of this 24k-node graph is SLOW; expect
        # 5+ minutes the first run. Subsequent loads are fast.)
        providers = ["CPUExecutionProvider"]
    return ort.InferenceSession(str(onnx_path), providers=providers)


def separate(mix: np.ndarray, sample_rate: int,
             onnx_path: Path = DEFAULT_ONNX,
             providers: list[str] | None = None,
             verbose: bool = True) -> np.ndarray:
    """Run chunked overlap-add separation on a full-length mix.

    Args:
        mix: (channels, samples) float32 in [-1, 1]. Must be 2-channel stereo
             at SAMPLE_RATE. Resample first with `soxr`, `librosa`, or `ffmpeg`
             if needed.
        sample_rate: must equal SAMPLE_RATE.
        onnx_path: path to htdemucs_ft_drums.onnx (defaults to alongside this
                   script).
        providers: list of onnxruntime execution providers; defaults to CPU.
        verbose: print per-chunk progress.

    Returns:
        (n_sources, channels, samples) float32 array. Only the row indexed by
        SOURCES.index(SPECIALIST_STEM) is meaningfully predicted — the other
        three rows are weakly-predicted by-products of the drum specialist
        and should not be used in production. For high-quality bass / vocals
        / other, ship the respective specialist ONNX file instead.
    """
    if sample_rate != SAMPLE_RATE:
        raise ValueError(
            f"This model is bound to {SAMPLE_RATE} Hz; got {sample_rate}. "
            "Resample your input first.")
    if mix.ndim != 2 or mix.shape[0] != N_CHANNELS:
        raise ValueError(f"Expected (2, samples) input, got {mix.shape}")

    sess = _load_session(onnx_path, providers)

    total_len = mix.shape[1]
    overlap = N_SAMPLES // 4
    stride = N_SAMPLES - overlap
    n_chunks = max(1, (total_len + stride - 1) // stride)

    if verbose:
        print(f"  input:    {total_len:,} samples ({total_len / sample_rate:.1f}s)")
        print(f"  segment:  {N_SAMPLES:,} samples ({SEGMENT_S}s)")
        print(f"  chunks:   {n_chunks}, stride {stride / sample_rate:.2f}s, "
              f"overlap {overlap / sample_rate:.2f}s")
        print(f"  provider: {sess.get_providers()[0]}")

    window = _make_transition_window(N_SAMPLES)
    out = np.zeros((len(SOURCES), N_CHANNELS, total_len), dtype=np.float32)
    weight = np.zeros(total_len, dtype=np.float32)

    t0 = time.perf_counter()
    for i in range(n_chunks):
        start = i * stride
        end = min(start + N_SAMPLES, total_len)
        chunk = mix[:, start:end]
        if chunk.shape[1] < N_SAMPLES:
            chunk = np.pad(chunk, ((0, 0), (0, N_SAMPLES - chunk.shape[1])),
                           mode="constant")
        x = chunk[np.newaxis, ...].astype(np.float32)  # (1, 2, N)
        stems = sess.run(["stems"], {"mix": x})[0][0]  # (4, 2, N)
        chunk_len = end - start
        w = window[:chunk_len]
        out[:, :, start:end] += stems[:, :, :chunk_len] * w
        weight[start:end] += w
        if verbose:
            print(f"    chunk {i+1}/{n_chunks}: "
                  f"{time.perf_counter() - t0:.1f}s elapsed")

    weight = np.maximum(weight, 1e-8)
    out /= weight
    if verbose:
        rtf = (time.perf_counter() - t0) / (total_len / sample_rate)
        print(f"  total:    {time.perf_counter() - t0:.2f}s (RTF {rtf:.2f})")
    return out


def separate_drums(input_path: str, onnx_path: Path = DEFAULT_ONNX,
                   providers: list[str] | None = None) -> np.ndarray:
    """Convenience: load audio, run separation, return only the drums stem.

    Returns: (channels, samples) float32 at 44.1 kHz.
    """
    audio, sr = sf.read(input_path, dtype="float32", always_2d=True)
    audio = audio.T
    if audio.shape[0] == 1:
        audio = np.tile(audio, (2, 1))
    elif audio.shape[0] > 2:
        audio = audio[:2]
    stems = separate(audio, sr, onnx_path=onnx_path, providers=providers)
    return stems[SOURCES.index(SPECIALIST_STEM)]


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("input", type=Path, help="Input audio file (wav/mp3/flac).")
    ap.add_argument("out_dir", type=Path, help="Output directory for .wav stems.")
    ap.add_argument("--onnx", type=Path, default=DEFAULT_ONNX,
                    help=f"Path to the .onnx file (default: {DEFAULT_ONNX})")
    ap.add_argument("--providers", type=str, default="cpu",
                    choices=["cpu", "coreml", "cuda", "dml"],
                    help="Execution provider (cpu is default and most portable).")
    ap.add_argument("--write-all-stems", action="store_true",
                    help="Also write bass/other/vocals (by-products, low quality).")
    args = ap.parse_args()

    providers_map = {
        "cpu": ["CPUExecutionProvider"],
        "coreml": ["CoreMLExecutionProvider", "CPUExecutionProvider"],
        "cuda": ["CUDAExecutionProvider", "CPUExecutionProvider"],
        "dml": ["DmlExecutionProvider", "CPUExecutionProvider"],
    }
    providers = providers_map[args.providers]
    args.out_dir.mkdir(parents=True, exist_ok=True)

    print(f"Loading {args.input} ...")
    audio, sr = sf.read(str(args.input), dtype="float32", always_2d=True)
    audio = audio.T
    if audio.shape[0] == 1:
        audio = np.tile(audio, (2, 1))
    elif audio.shape[0] > 2:
        audio = audio[:2]
    print(f"  shape {audio.shape}, sr {sr}")

    print(f"\nLoading {args.onnx} ({args.onnx.stat().st_size / 1e6:.1f} MB) ...")
    stems = separate(audio, sr, onnx_path=args.onnx, providers=providers)

    if args.write_all_stems:
        for i, src in enumerate(SOURCES):
            sf.write(str(args.out_dir / f"{src}.wav"), stems[i].T, sr)
            print(f"  wrote {args.out_dir / f'{src}.wav'}  "
                  f"{'*' if src == SPECIALIST_STEM else ''}")
        print("\n  * = meaningful prediction (specialist target)")
    else:
        drums = stems[SOURCES.index(SPECIALIST_STEM)]
        out_path = args.out_dir / "drums.wav"
        sf.write(str(out_path), drums.T, sr)
        print(f"  wrote {out_path}")


if __name__ == "__main__":
    main()