stem-splitter 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- stem_splitter/band_split_roformer.py +418 -0
- stem_splitter/cli.py +4 -0
- stem_splitter/inference.py +470 -0
- stem_splitter/transformer.py +257 -0
- stem_splitter-0.0.1.dist-info/METADATA +15 -0
- stem_splitter-0.0.1.dist-info/RECORD +9 -0
- stem_splitter-0.0.1.dist-info/WHEEL +5 -0
- stem_splitter-0.0.1.dist-info/entry_points.txt +2 -0
- stem_splitter-0.0.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,418 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
import torch.utils.checkpoint
|
|
4
|
+
import librosa
|
|
5
|
+
import numpy as np
|
|
6
|
+
import einops
|
|
7
|
+
from .transformer import Transformer, RMSNorm
|
|
8
|
+
|
|
9
|
+
from typing import List, Tuple
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def checkpoint_bypass(f, *args):
|
|
13
|
+
return f(*args)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def build_mel_band_indices(
|
|
17
|
+
sample_rate: int,
|
|
18
|
+
n_fft: int,
|
|
19
|
+
n_bands: int,
|
|
20
|
+
) -> list[np.ndarray[int]]:
|
|
21
|
+
"""
|
|
22
|
+
各 Mel バンドがカバーする「周波数ビン番号だけ」を返す。
|
|
23
|
+
返り値: list[np.ndarray] (len = n_bands)
|
|
24
|
+
"""
|
|
25
|
+
# --- Mel フィルタ → 0/1 マスク --------------------------------
|
|
26
|
+
mel_fb = librosa.filters.mel(sr=sample_rate, n_fft=n_fft, n_mels=n_bands)
|
|
27
|
+
mel_fb[0, 0] = 1.0
|
|
28
|
+
mel_fb[-1, -1] = 1.0
|
|
29
|
+
mask = mel_fb > 0 # (B, F) bool
|
|
30
|
+
|
|
31
|
+
return [np.where(mask[b])[0].astype(np.int32) for b in range(n_bands)]
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def build_band_indices(sample_rate: int = 44_100, n_fft: int = 2_048) -> Tuple[np.ndarray, List[np.ndarray]]:
|
|
35
|
+
"""
|
|
36
|
+
BSRoformer 用バンド分割 (各バンドの bin 数が必ず閾値以上)。
|
|
37
|
+
|
|
38
|
+
Returns
|
|
39
|
+
-------
|
|
40
|
+
freqs : np.ndarray
|
|
41
|
+
rfft の周波数 (長さ n_fft//2 + 1) [Hz]
|
|
42
|
+
bands : List[np.ndarray]
|
|
43
|
+
各要素が「そのバンドに含まれるビン番号」の 1‑D 配列
|
|
44
|
+
"""
|
|
45
|
+
freqs = np.fft.rfftfreq(n_fft, d=1 / sample_rate)
|
|
46
|
+
|
|
47
|
+
# バンドが始まるビンの周波数から必要 bin 数を決定
|
|
48
|
+
def bins_per_band(freq_hz: float) -> int:
|
|
49
|
+
if freq_hz < 1_000: # 0‑1 kHz
|
|
50
|
+
return 2
|
|
51
|
+
elif freq_hz < 2_000: # 1‑2 kHz
|
|
52
|
+
return 4
|
|
53
|
+
elif freq_hz < 4_000: # 2‑4 kHz
|
|
54
|
+
return 12
|
|
55
|
+
elif freq_hz < 8_000: # 4‑8 kHz
|
|
56
|
+
return 24
|
|
57
|
+
elif freq_hz < 16_000: # 8‑16 kHz
|
|
58
|
+
return 48
|
|
59
|
+
else: # ≥16 kHz → 後で 2 分割
|
|
60
|
+
return -1
|
|
61
|
+
|
|
62
|
+
freq_indices: List[np.ndarray] = []
|
|
63
|
+
idx = 0
|
|
64
|
+
n_bins = len(freqs)
|
|
65
|
+
|
|
66
|
+
while idx < n_bins:
|
|
67
|
+
# 16 kHz 以上は残りを 2 分割して終了
|
|
68
|
+
if freqs[idx] >= 16_000:
|
|
69
|
+
remaining = np.arange(idx, n_bins)
|
|
70
|
+
half = len(remaining) // 2
|
|
71
|
+
freq_indices.append(remaining[:half])
|
|
72
|
+
freq_indices.append(remaining[half:])
|
|
73
|
+
break
|
|
74
|
+
|
|
75
|
+
size = bins_per_band(freqs[idx])
|
|
76
|
+
band = np.arange(idx, min(idx + size, n_bins))
|
|
77
|
+
freq_indices.append(band)
|
|
78
|
+
idx += size # bin数 単位で次へ
|
|
79
|
+
|
|
80
|
+
return freqs, freq_indices
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
class BandSplit(torch.nn.Module):
|
|
84
|
+
def __init__(self, hidden_size: int, band_indices, num_channels: int, extra_windows: int):
|
|
85
|
+
super().__init__()
|
|
86
|
+
self.hidden_size = hidden_size
|
|
87
|
+
self.band_indices = band_indices
|
|
88
|
+
|
|
89
|
+
for i, idx in enumerate(band_indices):
|
|
90
|
+
self.register_buffer(f"band_idx_{i}", torch.tensor(idx, dtype=torch.long), persistent=False)
|
|
91
|
+
|
|
92
|
+
self.num_bands = len(band_indices)
|
|
93
|
+
self.to_features = torch.nn.ModuleList([])
|
|
94
|
+
for i in range(self.num_bands):
|
|
95
|
+
sub_band_freqs = len(band_indices[i]) * num_channels * 2 * (extra_windows + 1)
|
|
96
|
+
self.to_features.append(
|
|
97
|
+
torch.nn.Sequential(
|
|
98
|
+
RMSNorm(sub_band_freqs),
|
|
99
|
+
torch.nn.Linear(sub_band_freqs, hidden_size),
|
|
100
|
+
)
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
def forward(self, x):
|
|
104
|
+
# x: (B, T, C, F)
|
|
105
|
+
sub_band_list = []
|
|
106
|
+
for i, proj_layer in enumerate(self.to_features):
|
|
107
|
+
band_indices = getattr(self, f"band_idx_{i}")
|
|
108
|
+
# サブバンドインデックスで周波数軸から抜きだす
|
|
109
|
+
sub_band = x[..., band_indices] # (B, T, C, sub_band_freqs)
|
|
110
|
+
sub_band = einops.rearrange(sub_band, "b t c f -> b t (f c)")
|
|
111
|
+
sub_band = proj_layer(sub_band) # (B, T, hidden_size)
|
|
112
|
+
sub_band_list.append(sub_band)
|
|
113
|
+
|
|
114
|
+
return torch.stack(sub_band_list, dim=-2) # (B, T, num_bands, hidden_size)
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def band_mask_mlp(input_dim, band_dim, hidden_dim, depth):
|
|
118
|
+
layers = []
|
|
119
|
+
# 入力→隠れ→…→隠れ→(band_dim*2) の構造
|
|
120
|
+
dims = [input_dim] + [hidden_dim] * depth + [band_dim * 2]
|
|
121
|
+
for i in range(len(dims) - 1):
|
|
122
|
+
layers.append(nn.Linear(dims[i], dims[i + 1]))
|
|
123
|
+
# 最終層の後に活性化は入れない
|
|
124
|
+
if i < len(dims) - 2:
|
|
125
|
+
layers.append(nn.Tanh())
|
|
126
|
+
return nn.Sequential(*layers)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
class MaskEstimator(nn.Module):
|
|
130
|
+
def __init__(
|
|
131
|
+
self,
|
|
132
|
+
input_dim: int,
|
|
133
|
+
band_indices,
|
|
134
|
+
num_channels: int,
|
|
135
|
+
mlp_expansion_factor: int = 4,
|
|
136
|
+
depth: int = 2,
|
|
137
|
+
):
|
|
138
|
+
super().__init__()
|
|
139
|
+
self.input_dim = input_dim
|
|
140
|
+
self.num_channels = num_channels
|
|
141
|
+
self.band_indices = band_indices
|
|
142
|
+
|
|
143
|
+
self.num_bands = len(band_indices)
|
|
144
|
+
self.to_freqs = torch.nn.ModuleList([])
|
|
145
|
+
for i in range(self.num_bands):
|
|
146
|
+
sub_band_freqs = len(band_indices[i]) * num_channels * 2
|
|
147
|
+
self.to_freqs.append(
|
|
148
|
+
nn.Sequential(
|
|
149
|
+
band_mask_mlp(
|
|
150
|
+
input_dim=input_dim,
|
|
151
|
+
band_dim=sub_band_freqs,
|
|
152
|
+
hidden_dim=input_dim * mlp_expansion_factor,
|
|
153
|
+
depth=depth,
|
|
154
|
+
),
|
|
155
|
+
nn.GLU(dim=-1),
|
|
156
|
+
)
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
def forward(self, x):
|
|
160
|
+
# x: (B, T, K, D)
|
|
161
|
+
out = []
|
|
162
|
+
sub_band_list = torch.unbind(x, dim=2) # list([B, T, D])
|
|
163
|
+
assert len(sub_band_list) == len(self.to_freqs)
|
|
164
|
+
|
|
165
|
+
for i, proj_layer in enumerate(self.to_freqs):
|
|
166
|
+
sub_band = proj_layer(sub_band_list[i]) # (B, T, sub_band_freqs)
|
|
167
|
+
sub_band = einops.rearrange(
|
|
168
|
+
sub_band,
|
|
169
|
+
"b t (f c) -> b t c f",
|
|
170
|
+
c=self.num_channels * 2,
|
|
171
|
+
f=len(self.band_indices[i]),
|
|
172
|
+
)
|
|
173
|
+
out.append(sub_band)
|
|
174
|
+
|
|
175
|
+
return torch.concat(out, dim=-1) # (B, T, C, F)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class BSRoformer(nn.Module):
|
|
179
|
+
def __init__(
|
|
180
|
+
self,
|
|
181
|
+
dim: int,
|
|
182
|
+
num_layers: int = 12,
|
|
183
|
+
sample_rate: int = 44100,
|
|
184
|
+
num_channels: int = 2,
|
|
185
|
+
head_dim: int = 64,
|
|
186
|
+
num_heads: int = 8,
|
|
187
|
+
n_fft: int = 2048,
|
|
188
|
+
hop_length: int = 512,
|
|
189
|
+
n_bands: int = 60,
|
|
190
|
+
num_stems: int = 2,
|
|
191
|
+
band_split_type: str = "bs",
|
|
192
|
+
ffn_hidden_size_factor: int = 4,
|
|
193
|
+
dropout: float = 0.0,
|
|
194
|
+
use_shared_bias: bool = True,
|
|
195
|
+
use_gradient_checkpoint=True,
|
|
196
|
+
):
|
|
197
|
+
super().__init__()
|
|
198
|
+
self.n_fft = n_fft
|
|
199
|
+
self.window_size = n_fft
|
|
200
|
+
self.hop_length = hop_length
|
|
201
|
+
self.num_channels = num_channels
|
|
202
|
+
self.num_stems = num_stems
|
|
203
|
+
self.band_split_type = band_split_type
|
|
204
|
+
self.use_gradient_checkpoint = use_gradient_checkpoint
|
|
205
|
+
|
|
206
|
+
if band_split_type == "bs":
|
|
207
|
+
_, self.band_indices = build_band_indices(sample_rate=sample_rate, n_fft=n_fft)
|
|
208
|
+
elif band_split_type == "mel":
|
|
209
|
+
self.band_indices = build_mel_band_indices(sampling_rate=sample_rate, n_fft=n_fft, n_bands=n_bands)
|
|
210
|
+
else:
|
|
211
|
+
raise NotImplementedError("サポートしていない band_split_type です")
|
|
212
|
+
|
|
213
|
+
self.register_buffer("hann_window", torch.hann_window(self.window_size), persistent=False)
|
|
214
|
+
self.num_bands = len(self.band_indices)
|
|
215
|
+
self.band_split = BandSplit(
|
|
216
|
+
hidden_size=dim,
|
|
217
|
+
band_indices=self.band_indices,
|
|
218
|
+
num_channels=num_channels,
|
|
219
|
+
extra_windows=0,
|
|
220
|
+
)
|
|
221
|
+
|
|
222
|
+
self.layers = nn.ModuleList([])
|
|
223
|
+
if use_shared_bias:
|
|
224
|
+
hidden_size = head_dim * num_heads
|
|
225
|
+
self.shared_qkv_bias = nn.Parameter(torch.zeros(hidden_size * 3)) # QKV
|
|
226
|
+
self.shared_out_bias = nn.Parameter(torch.zeros(dim)) # OUT
|
|
227
|
+
|
|
228
|
+
for _ in range(num_layers):
|
|
229
|
+
time_roformer = Transformer(
|
|
230
|
+
input_dim=dim,
|
|
231
|
+
head_dim=head_dim,
|
|
232
|
+
num_layers=1,
|
|
233
|
+
num_heads=num_heads,
|
|
234
|
+
ffn_hidden_size_factor=ffn_hidden_size_factor,
|
|
235
|
+
dropout=dropout,
|
|
236
|
+
shared_qkv_bias=self.shared_qkv_bias,
|
|
237
|
+
shared_out_bias=self.shared_out_bias,
|
|
238
|
+
)
|
|
239
|
+
band_roformer = Transformer(
|
|
240
|
+
input_dim=dim,
|
|
241
|
+
head_dim=head_dim,
|
|
242
|
+
num_layers=1,
|
|
243
|
+
num_heads=num_heads,
|
|
244
|
+
ffn_hidden_size_factor=ffn_hidden_size_factor,
|
|
245
|
+
dropout=dropout,
|
|
246
|
+
shared_qkv_bias=self.shared_qkv_bias,
|
|
247
|
+
shared_out_bias=self.shared_out_bias,
|
|
248
|
+
)
|
|
249
|
+
self.layers.append(nn.ModuleList([time_roformer, band_roformer]))
|
|
250
|
+
|
|
251
|
+
self.final_norm = RMSNorm(dim)
|
|
252
|
+
|
|
253
|
+
self.register_buffer(
|
|
254
|
+
"freq_indices",
|
|
255
|
+
torch.tensor(np.concatenate(self.band_indices), dtype=torch.long),
|
|
256
|
+
persistent=False,
|
|
257
|
+
)
|
|
258
|
+
# そのビンが何本のバンドに含まれるかを数える
|
|
259
|
+
counts = np.zeros(n_fft // 2 + 1, dtype=np.float32)
|
|
260
|
+
for idx in self.band_indices:
|
|
261
|
+
counts[idx] += 1
|
|
262
|
+
self.register_buffer(
|
|
263
|
+
"num_bands_per_freq",
|
|
264
|
+
torch.tensor(counts, dtype=torch.float32),
|
|
265
|
+
persistent=False,
|
|
266
|
+
) # shape (F_total,)
|
|
267
|
+
self.mask_estimators = nn.ModuleList(
|
|
268
|
+
[
|
|
269
|
+
MaskEstimator(
|
|
270
|
+
input_dim=dim,
|
|
271
|
+
band_indices=self.band_indices,
|
|
272
|
+
num_channels=num_channels,
|
|
273
|
+
depth=1,
|
|
274
|
+
mlp_expansion_factor=4,
|
|
275
|
+
)
|
|
276
|
+
for _ in range(num_stems)
|
|
277
|
+
]
|
|
278
|
+
)
|
|
279
|
+
|
|
280
|
+
def to_spectrogram(self, inputs: torch.Tensor) -> torch.Tensor:
|
|
281
|
+
# (B, C, T)
|
|
282
|
+
inputs = einops.rearrange(inputs, "b c t -> (b c) t")
|
|
283
|
+
|
|
284
|
+
windows: list[torch.Tensor] = [self.hann_window]
|
|
285
|
+
|
|
286
|
+
spectrograms = []
|
|
287
|
+
for win in windows:
|
|
288
|
+
spectrogram = torch.stft(
|
|
289
|
+
input=inputs,
|
|
290
|
+
n_fft=self.n_fft,
|
|
291
|
+
hop_length=self.hop_length,
|
|
292
|
+
win_length=self.window_size,
|
|
293
|
+
window=win.to(inputs.device),
|
|
294
|
+
center=True,
|
|
295
|
+
return_complex=True,
|
|
296
|
+
)
|
|
297
|
+
spectrogram = einops.rearrange(spectrogram, "(b c) f t -> b c f t", c=self.num_channels)
|
|
298
|
+
spectrograms.append(spectrogram) # list of [B, C, F, T]
|
|
299
|
+
|
|
300
|
+
spectrogram = torch.concat(spectrograms, dim=1)
|
|
301
|
+
spectrogram = torch.view_as_real(spectrogram)
|
|
302
|
+
spectrogram = einops.rearrange(spectrogram, "b c f t s -> b (c s) t f", s=2)
|
|
303
|
+
return spectrogram.to(inputs.dtype), spectrograms[0]
|
|
304
|
+
|
|
305
|
+
def mixture_consistency_projection(
|
|
306
|
+
self,
|
|
307
|
+
source_estimates: torch.Tensor, # [B, C, N, F, T] complex
|
|
308
|
+
mixture: torch.Tensor, # [B, C, F, T] complex
|
|
309
|
+
weights: torch.Tensor | None = None, # [B, C, N, F, T] real, 任意
|
|
310
|
+
) -> torch.Tensor:
|
|
311
|
+
residual = mixture[:, :, None] - source_estimates.sum(dim=2, keepdim=True)
|
|
312
|
+
if weights is None:
|
|
313
|
+
correction = residual / source_estimates.shape[2]
|
|
314
|
+
else:
|
|
315
|
+
weights = weights / (weights.sum(dim=2, keepdim=True) + 1e-8)
|
|
316
|
+
correction = weights * residual
|
|
317
|
+
return source_estimates + correction
|
|
318
|
+
|
|
319
|
+
def to_recon_audio(
|
|
320
|
+
self,
|
|
321
|
+
mask: torch.Tensor,
|
|
322
|
+
original_spec_complex: torch.Tensor,
|
|
323
|
+
original_length: int,
|
|
324
|
+
dtype: torch.dtype,
|
|
325
|
+
device: torch.device,
|
|
326
|
+
) -> torch.Tensor:
|
|
327
|
+
# mask: [B, C, N, F, T]
|
|
328
|
+
source_estimates = original_spec_complex[:, :, None] * mask
|
|
329
|
+
|
|
330
|
+
# mixture consistency
|
|
331
|
+
magnitude = source_estimates.abs().clamp_min(1e-8)
|
|
332
|
+
weights = magnitude**2
|
|
333
|
+
separated_spectrogram = self.mixture_consistency_projection(
|
|
334
|
+
source_estimates=source_estimates,
|
|
335
|
+
mixture=original_spec_complex,
|
|
336
|
+
weights=weights,
|
|
337
|
+
)
|
|
338
|
+
|
|
339
|
+
separated_spectrogram = einops.rearrange(separated_spectrogram, "b c n f t -> (b c n) f t")
|
|
340
|
+
recon_audio = torch.istft(
|
|
341
|
+
separated_spectrogram,
|
|
342
|
+
n_fft=self.window_size,
|
|
343
|
+
hop_length=self.hop_length,
|
|
344
|
+
win_length=self.window_size,
|
|
345
|
+
window=self.hann_window.to(device),
|
|
346
|
+
center=True,
|
|
347
|
+
return_complex=False,
|
|
348
|
+
length=original_length,
|
|
349
|
+
) # [B*C, T]
|
|
350
|
+
recon_audio = einops.rearrange(recon_audio, "(b c n) t -> b c n t", c=self.num_channels, n=self.num_stems)
|
|
351
|
+
recon_audio = einops.rearrange(recon_audio, "b c n t -> b n c t")
|
|
352
|
+
return recon_audio
|
|
353
|
+
|
|
354
|
+
def forward(self, x):
|
|
355
|
+
# x: (B, C, T)
|
|
356
|
+
if self.use_gradient_checkpoint or self.training:
|
|
357
|
+
checkpoint = torch.utils.checkpoint.checkpoint
|
|
358
|
+
else:
|
|
359
|
+
checkpoint = checkpoint_bypass
|
|
360
|
+
|
|
361
|
+
istft_length = x.shape[-1]
|
|
362
|
+
x, original_spec_complex = self.to_spectrogram(x) # x: [B, C*S, T, F]
|
|
363
|
+
x = einops.rearrange(x, "b c t f -> b t c f")
|
|
364
|
+
|
|
365
|
+
# 勾配が流れない問題に対処
|
|
366
|
+
x = checkpoint(self.band_split, x, use_reentrant=False) # (B, T, K, hidden_size)
|
|
367
|
+
|
|
368
|
+
for time_roformer, band_roformer in self.layers:
|
|
369
|
+
B, T, K, F = x.shape
|
|
370
|
+
# 時間軸Transformer
|
|
371
|
+
x = einops.rearrange(x, "b t k f -> (b k) t f") # [B*K, T, F]
|
|
372
|
+
x = checkpoint(time_roformer, x, use_reentrant=False)
|
|
373
|
+
x = einops.rearrange(x, "(b k) t f -> b t k f", k=K) # [B, T, K, F]
|
|
374
|
+
|
|
375
|
+
# バンド軸Transformer
|
|
376
|
+
x = x.reshape(B * T, K, F) # [B*T, K, F]
|
|
377
|
+
x = checkpoint(band_roformer, x, use_reentrant=False)
|
|
378
|
+
x = x.reshape(B, T, K, F)
|
|
379
|
+
|
|
380
|
+
x = self.final_norm(x)
|
|
381
|
+
|
|
382
|
+
# 音源分離マスク
|
|
383
|
+
mask_list = []
|
|
384
|
+
for mask_estimator in self.mask_estimators:
|
|
385
|
+
pred_mask = checkpoint(mask_estimator, x, use_reentrant=False) # [B, T, C*S, F]
|
|
386
|
+
pred_mask = einops.rearrange(pred_mask, "b t (c s) f -> b t c f s", s=2).contiguous()
|
|
387
|
+
pred_mask = torch.view_as_complex(pred_mask) # [B, T, C, F]
|
|
388
|
+
mask_list.append(pred_mask)
|
|
389
|
+
|
|
390
|
+
mask_all = torch.stack(mask_list, dim=-2) # [B, T, C, N, F]
|
|
391
|
+
if self.band_split_type == "mel":
|
|
392
|
+
# 周波数ビンで重なる部分を加算する
|
|
393
|
+
B, T, C, N, F_concat = mask_all.shape
|
|
394
|
+
F_total = self.n_fft // 2 + 1
|
|
395
|
+
mask_sum = torch.zeros((B, T, C, N, F_total), dtype=mask_all.dtype, device=mask_all.device)
|
|
396
|
+
freq_idx = self.freq_indices.view(1, 1, 1, 1, -1).expand(B, T, C, N, -1) # [B, T, C, N, F_concat]
|
|
397
|
+
mask_sum.scatter_add_(
|
|
398
|
+
dim=-1,
|
|
399
|
+
index=freq_idx,
|
|
400
|
+
src=mask_all,
|
|
401
|
+
)
|
|
402
|
+
|
|
403
|
+
# 重なった本数で割って平均
|
|
404
|
+
denom = self.num_bands_per_freq.clamp(min=1e-8) # [F_total]
|
|
405
|
+
denom = denom.view(1, 1, 1, 1, -1) # broadcast
|
|
406
|
+
mask_avg = mask_sum / denom # [B, T, C, N, F_total]
|
|
407
|
+
mask_avg = einops.rearrange(mask_avg, "b t c n f -> b c n f t")
|
|
408
|
+
else:
|
|
409
|
+
mask_avg = einops.rearrange(mask_all, "b t c n f -> b c n f t")
|
|
410
|
+
|
|
411
|
+
recon_audio = self.to_recon_audio(
|
|
412
|
+
mask=mask_avg,
|
|
413
|
+
original_spec_complex=original_spec_complex,
|
|
414
|
+
original_length=istft_length,
|
|
415
|
+
dtype=x.dtype,
|
|
416
|
+
device=x.device,
|
|
417
|
+
)
|
|
418
|
+
return recon_audio
|
stem_splitter/cli.py
ADDED
|
@@ -0,0 +1,470 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import os
|
|
5
|
+
import sys
|
|
6
|
+
import urllib.request
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
from typing import Dict, Optional, Tuple, Union
|
|
10
|
+
|
|
11
|
+
import librosa
|
|
12
|
+
import numpy as np
|
|
13
|
+
import math
|
|
14
|
+
import torch
|
|
15
|
+
import torchaudio
|
|
16
|
+
import warnings
|
|
17
|
+
from tqdm import tqdm
|
|
18
|
+
|
|
19
|
+
from .band_split_roformer import BSRoformer
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass
|
|
23
|
+
class SeparationConfig:
|
|
24
|
+
# 推論関連
|
|
25
|
+
target_sample_rate: int = 44100
|
|
26
|
+
device_preference: Optional[str] = None # "cuda" / "cpu" / "mps"
|
|
27
|
+
use_half_precision: bool = False
|
|
28
|
+
|
|
29
|
+
chunk_size: int = 588_800 # 約 13.35 秒 @ 44.1kHz
|
|
30
|
+
hop_size: Optional[int] = None # 既定: chunk_size // 2(50% overlap)
|
|
31
|
+
window_type: str = "hann" # "hann" 推奨
|
|
32
|
+
|
|
33
|
+
stem_names: Tuple[str, ...] = (
|
|
34
|
+
"bass",
|
|
35
|
+
"drums",
|
|
36
|
+
"other",
|
|
37
|
+
"vocals",
|
|
38
|
+
"guitar",
|
|
39
|
+
"piano",
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
model_name: str = "bs_roformer"
|
|
43
|
+
hf_repo_id: Optional[str] = None
|
|
44
|
+
hf_filename: Optional[str] = None
|
|
45
|
+
hf_revision: str = "main"
|
|
46
|
+
expected_sha256: Optional[str] = None # 任意:完全性検証に使用
|
|
47
|
+
|
|
48
|
+
# キャッシュ先(未指定ならユーザホーム配下)
|
|
49
|
+
cache_dir: Optional[Path] = None
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
# モデル名→Hugging Face 上の既定情報(必要に応じて上書き)
|
|
53
|
+
MODEL_REGISTRY: Dict[str, Dict[str, str]] = {
|
|
54
|
+
"bs_roformer": {
|
|
55
|
+
"repo_id": "anime-song/stem-splitter",
|
|
56
|
+
"filename": "stem_splitter.pt",
|
|
57
|
+
"revision": "main",
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _save_wav_no_warning(out_path: Path, samples: torch.Tensor, sample_rate: int) -> None:
|
|
63
|
+
"""
|
|
64
|
+
警告を出さずに WAV 保存する。
|
|
65
|
+
優先度: TorchCodec > torchaudio.save_with_torchcodec > torchaudio.save(最後の手段)
|
|
66
|
+
入力: samples (C, T), float32 [-1, 1] 推奨
|
|
67
|
+
"""
|
|
68
|
+
if samples.device.type != "cpu":
|
|
69
|
+
samples = samples.cpu()
|
|
70
|
+
samples = samples.contiguous().to(torch.float32)
|
|
71
|
+
|
|
72
|
+
# 可能なら TorchCodec を直接使用(推奨)
|
|
73
|
+
try:
|
|
74
|
+
from torchcodec.encoders import AudioEncoder # pip install torchcodec
|
|
75
|
+
|
|
76
|
+
AudioEncoder(samples, sample_rate=sample_rate).to_file(str(out_path))
|
|
77
|
+
return
|
|
78
|
+
except Exception:
|
|
79
|
+
pass
|
|
80
|
+
|
|
81
|
+
# torchaudio の TorchCodec 経由 API(存在する場合)
|
|
82
|
+
if hasattr(torchaudio, "save_with_torchcodec"):
|
|
83
|
+
torchaudio.save_with_torchcodec(str(out_path), samples, sample_rate, channels_first=True)
|
|
84
|
+
return
|
|
85
|
+
|
|
86
|
+
warnings.filterwarnings(
|
|
87
|
+
"ignore",
|
|
88
|
+
message="In 2.9, this function's implementation will be changed to use torchaudio.save_with_torchcodec",
|
|
89
|
+
category=UserWarning,
|
|
90
|
+
module=r"torchaudio\._backend\.utils",
|
|
91
|
+
)
|
|
92
|
+
warnings.filterwarnings(
|
|
93
|
+
"ignore",
|
|
94
|
+
message="StreamWriter has been deprecated",
|
|
95
|
+
category=UserWarning,
|
|
96
|
+
module=r"torchaudio\._backend\.ffmpeg",
|
|
97
|
+
)
|
|
98
|
+
torchaudio.save(str(out_path), samples, sample_rate)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def resolve_device(device_preference: Optional[str]) -> torch.device:
|
|
102
|
+
if device_preference is not None:
|
|
103
|
+
return torch.device(device_preference)
|
|
104
|
+
if torch.cuda.is_available():
|
|
105
|
+
return torch.device("cuda")
|
|
106
|
+
if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
|
|
107
|
+
return torch.device("mps")
|
|
108
|
+
return torch.device("cpu")
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def default_cache_dir() -> Path:
|
|
112
|
+
return Path(os.environ.get("STEM_SPLITTER_CACHE", Path.home() / ".cache" / "stem_splitter" / "weights"))
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def get_model_info(config: SeparationConfig) -> Tuple[str, str, str]:
|
|
116
|
+
"""
|
|
117
|
+
config から Hugging Face の repo / filename / revision を解決。
|
|
118
|
+
未指定なら MODEL_REGISTRY を参照。
|
|
119
|
+
"""
|
|
120
|
+
repo_id = config.hf_repo_id
|
|
121
|
+
filename = config.hf_filename
|
|
122
|
+
revision = config.hf_revision
|
|
123
|
+
|
|
124
|
+
if (repo_id is None or filename is None) and config.model_name in MODEL_REGISTRY:
|
|
125
|
+
fallback = MODEL_REGISTRY[config.model_name]
|
|
126
|
+
repo_id = repo_id or fallback["repo_id"]
|
|
127
|
+
filename = filename or fallback["filename"]
|
|
128
|
+
revision = revision or fallback.get("revision", "main")
|
|
129
|
+
|
|
130
|
+
if repo_id is None or filename is None:
|
|
131
|
+
raise ValueError(
|
|
132
|
+
"Hugging Face の重みファイル情報が不足しています。"
|
|
133
|
+
" SeparationConfig(hf_repo_id=..., hf_filename=...) を指定するか、MODEL_REGISTRY を更新してください。"
|
|
134
|
+
)
|
|
135
|
+
return repo_id, filename, revision
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def build_weight_path(config: SeparationConfig, filename: str) -> Path:
|
|
139
|
+
cache_root = config.cache_dir or default_cache_dir()
|
|
140
|
+
cache_root.mkdir(parents=True, exist_ok=True)
|
|
141
|
+
return cache_root / filename
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
def sha256sum(file_path: Path) -> str:
|
|
145
|
+
h = hashlib.sha256()
|
|
146
|
+
with file_path.open("rb") as f:
|
|
147
|
+
for chunk in iter(lambda: f.read(8192), b""):
|
|
148
|
+
h.update(chunk)
|
|
149
|
+
return h.hexdigest()
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def download_from_huggingface(
|
|
153
|
+
repo_id: str,
|
|
154
|
+
filename: str,
|
|
155
|
+
revision: str,
|
|
156
|
+
destination: Path,
|
|
157
|
+
) -> None:
|
|
158
|
+
base_url = f"https://huggingface.co/{repo_id}/resolve/{revision}/{filename}?download=true"
|
|
159
|
+
|
|
160
|
+
with urllib.request.urlopen(base_url) as response:
|
|
161
|
+
total_size = int(response.headers.get("Content-Length", "0"))
|
|
162
|
+
block_size = 1024 * 1024 # 1MB
|
|
163
|
+
tmp_path = destination.with_suffix(".tmp")
|
|
164
|
+
|
|
165
|
+
with (
|
|
166
|
+
tmp_path.open("wb") as out_file,
|
|
167
|
+
tqdm(
|
|
168
|
+
total=total_size if total_size > 0 else None,
|
|
169
|
+
unit="B",
|
|
170
|
+
unit_scale=True,
|
|
171
|
+
desc=f"Downloading {filename}",
|
|
172
|
+
) as progress_bar,
|
|
173
|
+
):
|
|
174
|
+
while True:
|
|
175
|
+
data = response.read(block_size)
|
|
176
|
+
if not data:
|
|
177
|
+
break
|
|
178
|
+
out_file.write(data)
|
|
179
|
+
progress_bar.update(len(data))
|
|
180
|
+
|
|
181
|
+
tmp_path.replace(destination)
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def ensure_weight_file(config: SeparationConfig) -> Path:
|
|
185
|
+
repo_id, filename, revision = get_model_info(config)
|
|
186
|
+
weight_path = build_weight_path(config, filename)
|
|
187
|
+
|
|
188
|
+
if weight_path.exists():
|
|
189
|
+
if config.expected_sha256:
|
|
190
|
+
actual = sha256sum(weight_path)
|
|
191
|
+
if actual.lower() != config.expected_sha256.lower():
|
|
192
|
+
print(
|
|
193
|
+
f"[WARN] 既存の重みファイルの SHA256 が一致しません。再取得します。\n"
|
|
194
|
+
f" expected={config.expected_sha256}\n actual ={actual}"
|
|
195
|
+
)
|
|
196
|
+
weight_path.unlink(missing_ok=True)
|
|
197
|
+
else:
|
|
198
|
+
# 検証なしでそのまま使用
|
|
199
|
+
return weight_path
|
|
200
|
+
|
|
201
|
+
# ダウンロード
|
|
202
|
+
print(f"[INFO] 重みファイルが見つかりません。Hugging Face から取得します: {repo_id}/{filename}@{revision}")
|
|
203
|
+
download_from_huggingface(repo_id, filename, revision, weight_path)
|
|
204
|
+
|
|
205
|
+
if config.expected_sha256:
|
|
206
|
+
actual = sha256sum(weight_path)
|
|
207
|
+
if actual.lower() != config.expected_sha256.lower():
|
|
208
|
+
weight_path.unlink(missing_ok=True)
|
|
209
|
+
raise RuntimeError("ダウンロードした重みファイルの SHA256 が一致しません。")
|
|
210
|
+
|
|
211
|
+
return weight_path
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def load_state_dict_flex(path: Path) -> Dict[str, torch.Tensor]:
|
|
215
|
+
"""
|
|
216
|
+
checkpoint 形式の違いに寛容なロード(state_dict 直/ラップの両対応)。
|
|
217
|
+
"""
|
|
218
|
+
try:
|
|
219
|
+
checkpoint = torch.load(str(path), map_location="cpu", weights_only=False)
|
|
220
|
+
except TypeError:
|
|
221
|
+
# PyTorch の引数差異対策(古い/新しい両対応)
|
|
222
|
+
checkpoint = torch.load(str(path), map_location="cpu")
|
|
223
|
+
|
|
224
|
+
if isinstance(checkpoint, dict):
|
|
225
|
+
# 典型例: {"state_dict": ..., ...} / {"model": ..., ...} / そのまま state_dict
|
|
226
|
+
for key in ("state_dict", "model", "ema_state_dict"):
|
|
227
|
+
if key in checkpoint and isinstance(checkpoint[key], dict):
|
|
228
|
+
return checkpoint[key]
|
|
229
|
+
return checkpoint # そのまま state_dict とみなす
|
|
230
|
+
raise RuntimeError("未知のチェックポイント形式です。dict ではありません。")
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
def load_mss_model(config: SeparationConfig, device: torch.device) -> torch.nn.Module:
|
|
234
|
+
"""
|
|
235
|
+
BSRoformer を作成し、重みを読み込んで返します。
|
|
236
|
+
"""
|
|
237
|
+
weight_path = ensure_weight_file(config)
|
|
238
|
+
|
|
239
|
+
model = BSRoformer(
|
|
240
|
+
dim=256,
|
|
241
|
+
num_layers=12,
|
|
242
|
+
sample_rate=config.target_sample_rate,
|
|
243
|
+
num_channels=2,
|
|
244
|
+
head_dim=64,
|
|
245
|
+
num_heads=8,
|
|
246
|
+
n_fft=2048,
|
|
247
|
+
hop_length=512,
|
|
248
|
+
num_stems=len(config.stem_names),
|
|
249
|
+
)
|
|
250
|
+
|
|
251
|
+
# 重みを読み込み
|
|
252
|
+
state_dict = load_state_dict_flex(weight_path)
|
|
253
|
+
missing, unexpected = model.load_state_dict(state_dict, strict=False)
|
|
254
|
+
if missing:
|
|
255
|
+
print(f"[WARN] 欠落している重みキー: {sorted(missing)[:5]}{' ...' if len(missing) > 5 else ''}")
|
|
256
|
+
if unexpected:
|
|
257
|
+
print(f"[WARN] 予期しない重みキー: {sorted(unexpected)[:5]}{' ...' if len(unexpected) > 5 else ''}")
|
|
258
|
+
|
|
259
|
+
model.to(device)
|
|
260
|
+
if config.use_half_precision and device.type == "cuda":
|
|
261
|
+
model.half()
|
|
262
|
+
model.eval()
|
|
263
|
+
return model
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _separate_one_file(
|
|
267
|
+
input_wav_path: Path,
|
|
268
|
+
output_directory: Path,
|
|
269
|
+
config: SeparationConfig,
|
|
270
|
+
model: torch.nn.Module,
|
|
271
|
+
device: torch.device,
|
|
272
|
+
dtype: torch.dtype,
|
|
273
|
+
) -> Dict[str, Path]:
|
|
274
|
+
"""
|
|
275
|
+
単一の WAV を処理し、ステムごとの出力パスを返す。
|
|
276
|
+
入力: (C, T) / モデル入出力: (B, C, T) -> (B, N, C, T)
|
|
277
|
+
OLA: chunk_size, hop_size, Hann 窓(sum-of-windows 正規化)
|
|
278
|
+
"""
|
|
279
|
+
# 読み込み -> (C, T) に正規化
|
|
280
|
+
y, _ = librosa.load(str(input_wav_path), sr=config.target_sample_rate, mono=False)
|
|
281
|
+
if y.ndim == 1:
|
|
282
|
+
y = y[None, :] # (1, T)
|
|
283
|
+
elif y.ndim == 2 and y.shape[0] > y.shape[1]:
|
|
284
|
+
# librosa は (T, C) を返すことがあるので (C, T) へ
|
|
285
|
+
y = y.T
|
|
286
|
+
y = y.astype(np.float32, copy=False) # (C, T)
|
|
287
|
+
|
|
288
|
+
channels, total_length = y.shape
|
|
289
|
+
|
|
290
|
+
# チャンク条件
|
|
291
|
+
chunk_size = int(config.chunk_size)
|
|
292
|
+
hop_size = int(config.hop_size) if config.hop_size is not None else chunk_size // 2
|
|
293
|
+
if hop_size <= 0 or hop_size > chunk_size:
|
|
294
|
+
raise ValueError("hop_size は 1..chunk_size の範囲で指定してください")
|
|
295
|
+
|
|
296
|
+
# padding 後の長さ(最後が半端でも必ず 1 チャンクぶん確保)
|
|
297
|
+
if total_length <= chunk_size:
|
|
298
|
+
padded_length = chunk_size
|
|
299
|
+
else:
|
|
300
|
+
steps = math.ceil((total_length - chunk_size) / hop_size)
|
|
301
|
+
padded_length = steps * hop_size + chunk_size
|
|
302
|
+
|
|
303
|
+
if padded_length > total_length:
|
|
304
|
+
pad_amount = padded_length - total_length
|
|
305
|
+
y = np.pad(y, ((0, 0), (0, pad_amount)), mode="constant")
|
|
306
|
+
|
|
307
|
+
# 窓・蓄積(sum-of-windows で正規化)
|
|
308
|
+
if config.window_type.lower() != "hann":
|
|
309
|
+
raise ValueError(f"未対応の window_type: {config.window_type}")
|
|
310
|
+
base_window = torch.hann_window(chunk_size, periodic=False, dtype=dtype, device=device)
|
|
311
|
+
|
|
312
|
+
num_stems = len(config.stem_names)
|
|
313
|
+
accum = np.zeros((num_stems, channels, padded_length), dtype=np.float32)
|
|
314
|
+
weight_sum = np.zeros(padded_length, dtype=np.float32) # ← 窓の和を積む
|
|
315
|
+
|
|
316
|
+
# チャンク推論ループ
|
|
317
|
+
for start in range(0, padded_length - chunk_size + 1, hop_size):
|
|
318
|
+
end = start + chunk_size
|
|
319
|
+
|
|
320
|
+
# 末尾の短い区間はここでゼロパディング(上の pad と二重になっても影響なし)
|
|
321
|
+
input_chunk_np = y[:, start:end]
|
|
322
|
+
if input_chunk_np.shape[1] < chunk_size:
|
|
323
|
+
pad = chunk_size - input_chunk_np.shape[1]
|
|
324
|
+
input_chunk_np = np.pad(input_chunk_np, ((0, 0), (0, pad)), mode="constant")
|
|
325
|
+
|
|
326
|
+
# (1, C, T)
|
|
327
|
+
input_chunk = torch.from_numpy(input_chunk_np).to(device=device, dtype=dtype).unsqueeze(0)
|
|
328
|
+
|
|
329
|
+
with torch.no_grad():
|
|
330
|
+
output_chunk = model(input_chunk) # (1, N, C, T_out)
|
|
331
|
+
|
|
332
|
+
if not isinstance(output_chunk, torch.Tensor) or output_chunk.ndim != 4:
|
|
333
|
+
raise RuntimeError("モデル出力は (B, N, C, T) の Tensor を想定しています。")
|
|
334
|
+
|
|
335
|
+
_, _, _, t_out = output_chunk.shape
|
|
336
|
+
window = (
|
|
337
|
+
base_window if t_out == chunk_size else torch.hann_window(t_out, periodic=False, dtype=dtype, device=device)
|
|
338
|
+
)
|
|
339
|
+
|
|
340
|
+
# 出力に 1 回だけ窓掛け → 合成は窓の「和」で割る
|
|
341
|
+
windowed = (output_chunk * window.view(1, 1, 1, -1)).squeeze(0) # (N, C, T)
|
|
342
|
+
out_np = windowed.to(torch.float32).cpu().numpy()
|
|
343
|
+
|
|
344
|
+
accum[:, :, start : start + t_out] += out_np
|
|
345
|
+
weight_sum[start : start + t_out] += window.to(torch.float32).cpu().numpy()
|
|
346
|
+
|
|
347
|
+
del input_chunk, output_chunk, windowed # メモリ節約
|
|
348
|
+
|
|
349
|
+
# sum-of-windows で正規化(端も自然に補正される)
|
|
350
|
+
eps = 1e-8
|
|
351
|
+
weight_sum = np.maximum(weight_sum, eps)
|
|
352
|
+
accum /= weight_sum[None, None, :]
|
|
353
|
+
|
|
354
|
+
# 元の長さにトリムして保存
|
|
355
|
+
trim_length = total_length
|
|
356
|
+
base_name = input_wav_path.stem
|
|
357
|
+
saved_paths: Dict[str, Path] = {}
|
|
358
|
+
|
|
359
|
+
for i, stem_name in enumerate(config.stem_names):
|
|
360
|
+
stem_array = accum[i, :, :trim_length] # (C, T)
|
|
361
|
+
stem_tensor = torch.from_numpy(stem_array) # float32
|
|
362
|
+
out_path = output_directory / base_name / f"{base_name}_{stem_name}.wav"
|
|
363
|
+
out_path.parent.mkdir(parents=True, exist_ok=True)
|
|
364
|
+
_save_wav_no_warning(out_path, stem_tensor, config.target_sample_rate)
|
|
365
|
+
saved_paths[stem_name] = out_path
|
|
366
|
+
|
|
367
|
+
return saved_paths
|
|
368
|
+
|
|
369
|
+
|
|
370
|
+
def separate_stems(
|
|
371
|
+
input_audio_path: Union[str, Path],
|
|
372
|
+
output_directory: Union[str, Path],
|
|
373
|
+
config: Optional[SeparationConfig] = None,
|
|
374
|
+
) -> Union[Dict[str, Path], Dict[Path, Dict[str, Path]]]:
|
|
375
|
+
if config is None:
|
|
376
|
+
config = SeparationConfig()
|
|
377
|
+
|
|
378
|
+
input_path = Path(input_audio_path)
|
|
379
|
+
output_dir = Path(output_directory)
|
|
380
|
+
output_dir.mkdir(parents=True, exist_ok=True)
|
|
381
|
+
|
|
382
|
+
# デバイス・dtype・モデルはここで1回だけ
|
|
383
|
+
device = resolve_device(config.device_preference)
|
|
384
|
+
dtype = torch.float16 if (config.use_half_precision and device.type == "cuda") else torch.float32
|
|
385
|
+
model = load_mss_model(config, device=device)
|
|
386
|
+
|
|
387
|
+
if input_path.is_file():
|
|
388
|
+
if input_path.suffix.lower() != ".wav":
|
|
389
|
+
raise ValueError("WAV 以外は許可していません。入力は .wav にしてください。")
|
|
390
|
+
return _separate_one_file(input_path, output_dir, config, model, device, dtype)
|
|
391
|
+
|
|
392
|
+
if input_path.is_dir():
|
|
393
|
+
wav_files = sorted([p for p in input_path.rglob("*") if p.is_file() and p.suffix.lower() == ".wav"])
|
|
394
|
+
if not wav_files:
|
|
395
|
+
raise FileNotFoundError("指定ディレクトリに .wav ファイルが見つかりません。")
|
|
396
|
+
|
|
397
|
+
results: Dict[Path, Dict[str, Path]] = {}
|
|
398
|
+
for wav_path in wav_files:
|
|
399
|
+
current_out_dir = output_dir # フラットに保存する場合はこちら
|
|
400
|
+
result = _separate_one_file(wav_path, current_out_dir, config, model, device, dtype)
|
|
401
|
+
results[wav_path] = result
|
|
402
|
+
return results
|
|
403
|
+
|
|
404
|
+
raise FileNotFoundError(f"入力パスが存在しません: {input_path}")
|
|
405
|
+
|
|
406
|
+
|
|
407
|
+
def _build_arg_parser():
|
|
408
|
+
import argparse
|
|
409
|
+
|
|
410
|
+
parser = argparse.ArgumentParser(description="Stem Splitter Inference (BSRoformer, OLA)")
|
|
411
|
+
parser.add_argument("input_audio_path", type=Path, help="入力音声ファイルのパス")
|
|
412
|
+
parser.add_argument("--out-dir", type=Path, required=True, help="出力ディレクトリ")
|
|
413
|
+
parser.add_argument("--device", default=None, help="cuda / cpu / mps(未指定で自動判定)")
|
|
414
|
+
parser.add_argument("--sr", type=int, default=44100, help="サンプルレート")
|
|
415
|
+
parser.add_argument("--no-half", action="store_true", help="半精度を無効化(既定はCUDA上で有効)")
|
|
416
|
+
|
|
417
|
+
parser.add_argument("--chunk-size", type=int, default=588_800, help="チャンク長(サンプル数)")
|
|
418
|
+
parser.add_argument("--hop-size", type=int, default=None, help="ホップ長(未指定で chunk_size//2)")
|
|
419
|
+
parser.add_argument("--window", type=str, default="hann", help="ウィンドウ種別(hann のみ対応)")
|
|
420
|
+
|
|
421
|
+
# 重み指定の上書き
|
|
422
|
+
parser.add_argument("--model-name", default="bs_roformer", help="MODEL_REGISTRY のキー")
|
|
423
|
+
parser.add_argument("--hf-repo", default=None, help="Hugging Face repo_id")
|
|
424
|
+
parser.add_argument("--hf-file", default=None, help="Hugging Face ファイル名")
|
|
425
|
+
parser.add_argument("--hf-rev", default="main", help="Hugging Face リビジョン")
|
|
426
|
+
parser.add_argument("--weights-cache", type=Path, default=None, help="重みキャッシュディレクトリ")
|
|
427
|
+
parser.add_argument("--sha256", default=None, help="重みファイル SHA256(任意)")
|
|
428
|
+
return parser
|
|
429
|
+
|
|
430
|
+
|
|
431
|
+
def main():
|
|
432
|
+
parser = _build_arg_parser()
|
|
433
|
+
args = parser.parse_args()
|
|
434
|
+
|
|
435
|
+
config = SeparationConfig(
|
|
436
|
+
target_sample_rate=args.sr,
|
|
437
|
+
device_preference=args.device,
|
|
438
|
+
use_half_precision=not args.no_half,
|
|
439
|
+
chunk_size=args.chunk_size,
|
|
440
|
+
hop_size=args.hop_size,
|
|
441
|
+
window_type=args.window,
|
|
442
|
+
model_name=args.model_name,
|
|
443
|
+
hf_repo_id=args.hf_repo,
|
|
444
|
+
hf_filename=args.hf_file,
|
|
445
|
+
hf_revision=args.hf_rev,
|
|
446
|
+
expected_sha256=args.sha256,
|
|
447
|
+
cache_dir=args.weights_cache,
|
|
448
|
+
)
|
|
449
|
+
|
|
450
|
+
try:
|
|
451
|
+
result = separate_stems(args.input_audio_path, args.out_dir, config)
|
|
452
|
+
except Exception as exc:
|
|
453
|
+
print(f"[ERROR] 推論に失敗しました: {exc}", file=sys.stderr)
|
|
454
|
+
sys.exit(1)
|
|
455
|
+
|
|
456
|
+
# 表示
|
|
457
|
+
if isinstance(result, dict) and result and isinstance(next(iter(result.values())), Path):
|
|
458
|
+
# 単一ファイルケース: {stem: path}
|
|
459
|
+
for stem_name, out_path in result.items():
|
|
460
|
+
print(f"{stem_name}: {out_path}")
|
|
461
|
+
else:
|
|
462
|
+
# ディレクトリケース: {input_wav: {stem: path}}
|
|
463
|
+
for input_wav, stems_dict in result.items():
|
|
464
|
+
print(f"[{input_wav}]")
|
|
465
|
+
for stem_name, out_path in stems_dict.items():
|
|
466
|
+
print(f" {stem_name}: {out_path}")
|
|
467
|
+
|
|
468
|
+
|
|
469
|
+
if __name__ == "__main__":
|
|
470
|
+
main()
|
|
@@ -0,0 +1,257 @@
|
|
|
1
|
+
import torch.nn as nn
|
|
2
|
+
import torch.nn.functional as F
|
|
3
|
+
import torch
|
|
4
|
+
from torch.nn.attention import SDPBackend, sdpa_kernel
|
|
5
|
+
import einops
|
|
6
|
+
from typing import Tuple
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def choose_low_precision_dtype() -> torch.dtype:
|
|
10
|
+
"""
|
|
11
|
+
GPU の機能を調べて BF16 → FP16 → FP32 の順に
|
|
12
|
+
最も高速な演算 dtype を返す。
|
|
13
|
+
"""
|
|
14
|
+
if not torch.cuda.is_available():
|
|
15
|
+
return torch.float32 # CPU 実行なら FP32 一択
|
|
16
|
+
|
|
17
|
+
# Ampere (sm80) 以降ならほぼ BF16 演算に対応
|
|
18
|
+
if torch.cuda.is_bf16_supported(): # PyTorch 2.1+
|
|
19
|
+
return torch.bfloat16
|
|
20
|
+
|
|
21
|
+
major_cc, _ = torch.cuda.get_device_capability()
|
|
22
|
+
# Pascal (sm60) 以降なら FP16 演算ユニットあり
|
|
23
|
+
if major_cc >= 6:
|
|
24
|
+
return torch.float16
|
|
25
|
+
|
|
26
|
+
return torch.float32 # それ以前の Maxwell など
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class RMSNorm(nn.Module):
|
|
30
|
+
def __init__(self, dim, eps: float = 5.960464477539063e-08): # 0x1p-24
|
|
31
|
+
super().__init__()
|
|
32
|
+
self.scale = dim**0.5
|
|
33
|
+
self.gamma = nn.Parameter(torch.ones(dim))
|
|
34
|
+
self.eps = eps
|
|
35
|
+
|
|
36
|
+
def forward(self, x):
|
|
37
|
+
l2_norm = torch.linalg.norm(x, dim=-1, keepdim=True)
|
|
38
|
+
denom = torch.maximum(l2_norm, torch.full_like(l2_norm, self.eps))
|
|
39
|
+
normalized_x = x / denom
|
|
40
|
+
return normalized_x * self.scale * self.gamma
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class RotaryEmbeddings(torch.nn.Module):
|
|
44
|
+
"""
|
|
45
|
+
RoPE 用の sin・cos テーブルをキャッシュし、必要に応じて伸張/切り詰めるクラス。
|
|
46
|
+
|
|
47
|
+
Args:
|
|
48
|
+
head_dim: 1 ヘッドあたりの埋め込み次元 (必ず偶数にする)
|
|
49
|
+
max_seq_len: 事前に準備しておく最大シーケンス長
|
|
50
|
+
base_theta: 周波数スケーリング係数 (多くの論文では 10000.0)
|
|
51
|
+
learned: True にすると sin, cos をパラメータとして学習させられる
|
|
52
|
+
"""
|
|
53
|
+
|
|
54
|
+
def __init__(
|
|
55
|
+
self,
|
|
56
|
+
head_dim: int,
|
|
57
|
+
max_seq_len: int = 2048,
|
|
58
|
+
base_theta: float = 10000.0,
|
|
59
|
+
learned: bool = False,
|
|
60
|
+
device: torch.device | None = None,
|
|
61
|
+
):
|
|
62
|
+
super().__init__()
|
|
63
|
+
|
|
64
|
+
if head_dim % 2 != 0:
|
|
65
|
+
raise ValueError("head_dim (=1 ヘッドの次元数) は偶数にしてください。")
|
|
66
|
+
|
|
67
|
+
self.head_dim = head_dim
|
|
68
|
+
self.base_theta = base_theta
|
|
69
|
+
self.max_seq_len = max_seq_len
|
|
70
|
+
|
|
71
|
+
# 角周波数: θ_k = (θ_base)^(2k / d)
|
|
72
|
+
freqs = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device)
|
|
73
|
+
inv_freq = 1.0 / (base_theta ** (freqs / head_dim)) # (head_dim/2,)
|
|
74
|
+
|
|
75
|
+
# time 方向へアウター積 → (max_seq_len, head_dim/2)
|
|
76
|
+
t = torch.arange(max_seq_len, device=device).float()
|
|
77
|
+
sinusoid_inp = torch.einsum("i,j->ij", t, inv_freq)
|
|
78
|
+
|
|
79
|
+
sin, cos = (
|
|
80
|
+
sinusoid_inp.sin(),
|
|
81
|
+
sinusoid_inp.cos(),
|
|
82
|
+
) # 各が (max_seq_len, head_dim/2)
|
|
83
|
+
|
|
84
|
+
if learned:
|
|
85
|
+
self.register_parameter("sin_cached", torch.nn.Parameter(sin))
|
|
86
|
+
self.register_parameter("cos_cached", torch.nn.Parameter(cos))
|
|
87
|
+
else:
|
|
88
|
+
self.register_buffer("sin_cached", sin, persistent=False)
|
|
89
|
+
self.register_buffer("cos_cached", cos, persistent=False)
|
|
90
|
+
|
|
91
|
+
def forward(
|
|
92
|
+
self,
|
|
93
|
+
seq_len: int,
|
|
94
|
+
dtype: torch.dtype = torch.float32,
|
|
95
|
+
device: torch.device | None = None,
|
|
96
|
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
97
|
+
"""
|
|
98
|
+
指定長の sin, cos を返す。
|
|
99
|
+
|
|
100
|
+
Returns:
|
|
101
|
+
cos: (seq_len, head_dim/2)
|
|
102
|
+
sin: (seq_len, head_dim/2)
|
|
103
|
+
"""
|
|
104
|
+
if seq_len > self.max_seq_len:
|
|
105
|
+
raise ValueError(f"要求シーケンス長 {seq_len} は max_seq_len={self.max_seq_len} を超えています。")
|
|
106
|
+
cos = self.cos_cached[:seq_len].to(dtype=dtype, device=device)
|
|
107
|
+
sin = self.sin_cached[:seq_len].to(dtype=dtype, device=device)
|
|
108
|
+
return cos, sin
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def rotate_half(x: torch.Tensor) -> torch.Tensor:
|
|
112
|
+
"""
|
|
113
|
+
偶数次元のテンソルを (…, 2i, 2i+1) → (…, -2i+1, 2i) のように 90° 回転させる。
|
|
114
|
+
具体的には (x_even, x_odd) → (-x_odd, x_even)。
|
|
115
|
+
"""
|
|
116
|
+
x_even, x_odd = x[..., 0::2], x[..., 1::2]
|
|
117
|
+
return torch.stack((-x_odd, x_even), dim=-1).flatten(-2)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def apply_rotary_embedding(
|
|
121
|
+
query: torch.Tensor, # (batch, num_heads, seq_len, head_dim)
|
|
122
|
+
key: torch.Tensor, # (batch, num_heads, seq_len, head_dim)
|
|
123
|
+
cos: torch.Tensor, # (seq_len, head_dim/2)
|
|
124
|
+
sin: torch.Tensor, # (seq_len, head_dim/2)
|
|
125
|
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
126
|
+
cos = cos[None, None, :, :]
|
|
127
|
+
sin = sin[None, None, :, :]
|
|
128
|
+
cos = torch.repeat_interleave(cos, 2, dim=-1)
|
|
129
|
+
sin = torch.repeat_interleave(sin, 2, dim=-1)
|
|
130
|
+
|
|
131
|
+
q_rot = query * cos + rotate_half(query) * sin
|
|
132
|
+
k_rot = key * cos + rotate_half(key) * sin
|
|
133
|
+
return q_rot, k_rot
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
class FeedForward(nn.Module):
|
|
137
|
+
def __init__(self, dim, ffn_hidden_size_factor=4, dropout=0.0):
|
|
138
|
+
super().__init__()
|
|
139
|
+
dim_inner = int(dim * ffn_hidden_size_factor)
|
|
140
|
+
self.net = nn.Sequential(
|
|
141
|
+
RMSNorm(dim),
|
|
142
|
+
nn.Linear(dim, dim_inner),
|
|
143
|
+
nn.GELU(),
|
|
144
|
+
nn.Dropout(dropout),
|
|
145
|
+
nn.Linear(dim_inner, dim),
|
|
146
|
+
nn.Dropout(dropout),
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
def forward(self, x):
|
|
150
|
+
return self.net(x)
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
class MultiHeadAttention(nn.Module):
|
|
154
|
+
def __init__(
|
|
155
|
+
self,
|
|
156
|
+
input_dim,
|
|
157
|
+
num_heads=8,
|
|
158
|
+
head_dim=64,
|
|
159
|
+
shared_qkv_bias=None,
|
|
160
|
+
shared_out_bias=None,
|
|
161
|
+
dropout=0.0,
|
|
162
|
+
):
|
|
163
|
+
super().__init__()
|
|
164
|
+
self.hidden_size = head_dim * num_heads
|
|
165
|
+
self.num_heads = num_heads
|
|
166
|
+
self.head_dim = head_dim
|
|
167
|
+
|
|
168
|
+
self.norm = RMSNorm(input_dim)
|
|
169
|
+
self.to_qkv = nn.Linear(input_dim, self.hidden_size * 3, bias=(shared_qkv_bias is not None))
|
|
170
|
+
if shared_qkv_bias is not None:
|
|
171
|
+
self.to_qkv.bias = shared_qkv_bias
|
|
172
|
+
|
|
173
|
+
self.to_gates = nn.Linear(input_dim, num_heads)
|
|
174
|
+
self.to_out = nn.Sequential(
|
|
175
|
+
nn.Linear(self.hidden_size, input_dim, bias=(shared_out_bias is not None)),
|
|
176
|
+
nn.Dropout(dropout),
|
|
177
|
+
)
|
|
178
|
+
if shared_out_bias is not None:
|
|
179
|
+
self.to_out[0].bias = shared_out_bias
|
|
180
|
+
|
|
181
|
+
self.rope = RotaryEmbeddings(
|
|
182
|
+
head_dim=self.head_dim,
|
|
183
|
+
learned=False,
|
|
184
|
+
)
|
|
185
|
+
self.lowp_dtype = choose_low_precision_dtype()
|
|
186
|
+
|
|
187
|
+
def forward(self, x):
|
|
188
|
+
x = self.norm(x)
|
|
189
|
+
|
|
190
|
+
q, k, v = einops.rearrange(self.to_qkv(x), "b t (qkv h d) -> qkv b h t d", qkv=3, h=self.num_heads)
|
|
191
|
+
|
|
192
|
+
cos, sin = self.rope(q.shape[-2], dtype=q.dtype, device=q.device)
|
|
193
|
+
q, k = apply_rotary_embedding(q, k, cos, sin)
|
|
194
|
+
|
|
195
|
+
q = q.to(self.lowp_dtype)
|
|
196
|
+
k = k.to(self.lowp_dtype)
|
|
197
|
+
v = v.to(self.lowp_dtype)
|
|
198
|
+
with sdpa_kernel(
|
|
199
|
+
[
|
|
200
|
+
SDPBackend.FLASH_ATTENTION,
|
|
201
|
+
SDPBackend.EFFICIENT_ATTENTION,
|
|
202
|
+
SDPBackend.MATH,
|
|
203
|
+
]
|
|
204
|
+
):
|
|
205
|
+
fetched = F.scaled_dot_product_attention(q, k, v)
|
|
206
|
+
|
|
207
|
+
gates = self.to_gates(x)
|
|
208
|
+
gates = gates.sigmoid()
|
|
209
|
+
|
|
210
|
+
out = fetched.to(x.dtype) * einops.rearrange(gates, "b n h -> b h n 1")
|
|
211
|
+
out = einops.rearrange(out, "b h t d -> b t (h d)")
|
|
212
|
+
return self.to_out(out)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
class Transformer(nn.Module):
|
|
216
|
+
def __init__(
|
|
217
|
+
self,
|
|
218
|
+
input_dim: int,
|
|
219
|
+
head_dim: int,
|
|
220
|
+
num_heads: int,
|
|
221
|
+
num_layers: int,
|
|
222
|
+
ffn_hidden_size_factor: int = 4,
|
|
223
|
+
dropout: float = 0.0,
|
|
224
|
+
shared_qkv_bias=None,
|
|
225
|
+
shared_out_bias=None,
|
|
226
|
+
output_norm: bool = False,
|
|
227
|
+
):
|
|
228
|
+
super().__init__()
|
|
229
|
+
self.layers = nn.ModuleList([])
|
|
230
|
+
|
|
231
|
+
for _ in range(num_layers):
|
|
232
|
+
attention = MultiHeadAttention(
|
|
233
|
+
input_dim=input_dim,
|
|
234
|
+
head_dim=head_dim,
|
|
235
|
+
num_heads=num_heads,
|
|
236
|
+
shared_qkv_bias=shared_qkv_bias,
|
|
237
|
+
shared_out_bias=shared_out_bias,
|
|
238
|
+
)
|
|
239
|
+
self.layers.append(
|
|
240
|
+
nn.ModuleList(
|
|
241
|
+
[
|
|
242
|
+
attention,
|
|
243
|
+
FeedForward(dim=input_dim, ffn_hidden_size_factor=ffn_hidden_size_factor),
|
|
244
|
+
]
|
|
245
|
+
)
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
self.norm = RMSNorm(input_dim) if output_norm else nn.Identity()
|
|
249
|
+
|
|
250
|
+
def forward(self, x):
|
|
251
|
+
# x: [B, T, F]
|
|
252
|
+
for attention, ffn in self.layers:
|
|
253
|
+
x = attention(x) + x
|
|
254
|
+
x = ffn(x) + x
|
|
255
|
+
|
|
256
|
+
x = self.norm(x)
|
|
257
|
+
return x
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: stem-splitter
|
|
3
|
+
Version: 0.0.1
|
|
4
|
+
Summary: Simple, readable audio stem separation library
|
|
5
|
+
Author: anime-song
|
|
6
|
+
License: MIT
|
|
7
|
+
Requires-Python: >=3.10
|
|
8
|
+
Description-Content-Type: text/markdown
|
|
9
|
+
Requires-Dist: numpy>=1.24
|
|
10
|
+
Requires-Dist: tqdm>=4.66
|
|
11
|
+
Requires-Dist: librosa>=0.10
|
|
12
|
+
Requires-Dist: einops>=0.7
|
|
13
|
+
Requires-Dist: torch
|
|
14
|
+
Requires-Dist: torchaudio
|
|
15
|
+
Requires-Dist: torchcodec
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
stem_splitter/band_split_roformer.py,sha256=-Ni4F0DBitLhWjmFupjb46pb8UjE3zSEY8gMvcuUjjM,15810
|
|
2
|
+
stem_splitter/cli.py,sha256=afs4cil4jXAR35sO26RcPWfRLSs7VRwvOC9Fn65e1JY,67
|
|
3
|
+
stem_splitter/inference.py,sha256=aXVYczaXpWiKsHDQpMSY1lZtqsIyEI-nEX-W892TQRM,18396
|
|
4
|
+
stem_splitter/transformer.py,sha256=JHlqbMFBI_8zSXhyp5G8bpXAvdldhDvIWMLlg0m0ZT8,8696
|
|
5
|
+
stem_splitter-0.0.1.dist-info/METADATA,sha256=Izh4xvehfyAu-tL7tKbBUPuVsCcbtLpe871suBr3IHo,391
|
|
6
|
+
stem_splitter-0.0.1.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
|
|
7
|
+
stem_splitter-0.0.1.dist-info/entry_points.txt,sha256=rMaUzno9dQgUZT0GH6j1NTieeDjQIOW8Nrri-EDojYg,57
|
|
8
|
+
stem_splitter-0.0.1.dist-info/top_level.txt,sha256=d-BKlmkJUV_uluTwOpHhfFhwyzD_9hYPI4VCmxiElwU,14
|
|
9
|
+
stem_splitter-0.0.1.dist-info/RECORD,,
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
stem_splitter
|