[Added] Working node and CLI tool

This commit is contained in:
Salvador E. Tropea
2025-07-02 08:33:52 -03:00
parent b0a1e03326
commit e0b2bde974
29 changed files with 2375 additions and 0 deletions
+187
View File
@@ -0,0 +1,187 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
# Developed using:
# - Netron to inspect the topology
# - onnx2pytorch.ConvertModel to find mapping details
# - Gemini 2.5 Pro to analyze the network, write the code and debug it
# I never saw the original implementation, this is a reconstruction from Kim_Vocals_2.onnx
#
# Found geometries:
#
# dim_f | channels | stages | ~params
# -------|----------|--------|------------
# 3072 | 48 | 5 | 16 684 228
# 2560 | 48 | 5 | 14 763 012
# 2048 | 48 | 5 | 13 191 108
# 2048 | 32 | 5 | 7 420 548
# 2048 | 32 | 4 | 5 478 276
from torch import nn
class FrequencyBranch(nn.Module):
""" Frequency-domain branch with Linear -> BatchNorm2d -> ReLU sequences. """
def __init__(self, channels, freq_dim, hidden_dim, bn_eps):
super().__init__()
self.sequence = nn.Sequential(
nn.Linear(freq_dim, hidden_dim, bias=False),
# Using the standard, verified nn.BatchNorm2d
nn.BatchNorm2d(num_features=channels, eps=bn_eps),
nn.ReLU(True),
nn.Linear(hidden_dim, freq_dim, bias=False),
nn.BatchNorm2d(num_features=channels, eps=bn_eps),
nn.ReLU(True)
)
def forward(self, x):
# Relies on PyTorch's nn.Linear broadcasting over the first 3 dims of (B,C,T,F)
# And nn.BatchNorm2d operating on the C dimension of the 4D tensor.
return self.sequence(x)
class TimeBranch(nn.Module):
"""
Time-domain branch using 3x3 convolutions.
"""
def __init__(self, channels):
super().__init__()
self.sequence = nn.Sequential(
nn.Conv2d(channels, channels, kernel_size=3, padding='same', bias=True),
nn.ReLU(True),
nn.Conv2d(channels, channels, kernel_size=3, padding='same', bias=True),
nn.ReLU(True),
nn.Conv2d(channels, channels, kernel_size=3, padding='same', bias=True),
nn.ReLU(True)
)
def forward(self, x):
return self.sequence(x)
class TDF_Block(nn.Module):
""" The main processing block, combining time and frequency branches.
This is a sequential-residual block, as shown in the ONNX graph. """
def __init__(self, channels, freq_dim, hidden_dim, bn_eps):
super().__init__()
self.time_branch = TimeBranch(channels)
self.freq_branch = FrequencyBranch(channels, freq_dim, hidden_dim, bn_eps)
def forward(self, x):
# 1. The input 'x' goes through the time branch first.
time_out = self.time_branch(x)
# 2. The output of the time branch is then fed into the frequency branch.
freq_out = self.freq_branch(time_out)
# 3. The final result is a residual connection:
# Output of Time Branch + Output of Frequency Branch
return time_out + freq_out
class Transpose(nn.Module):
""" A simple nn.Module wrapper for the permute operation """
def __init__(self, dims):
super().__init__()
self.dims = dims
def forward(self, x):
return x.permute(self.dims)
class MDX_Net(nn.Module):
"""
The complete U-Net architecture.
Fully parametric for frequency bins, channels, and number of stages.
This version uses your elegant interlaced ModuleList design for a clean,
dynamic structure that correctly matches the ONNX graph order.
"""
def __init__(self, dim_f=3072, ch=48, num_stages=5):
super().__init__()
# Validate input
if num_stages < 1 or num_stages > 12:
raise ValueError(f"num_stages must be between 1 and 12, but got {num_stages}")
self.num_stages = num_stages
# Define shared BatchNorm parameters
BN_EPS = 9.999999747378752e-06
freq_hidden_dim = dim_f // 8
# Allow others to know our creation parameters
self.dim_f = dim_f
self.ch = ch
self.num_stages = num_stages
# --- Initial Chain (always exists) ---
self.initial_conv = nn.Conv2d(4, ch, 1, bias=True)
self.initial_relu = nn.ReLU(True)
self.initial_transpose = Transpose(dims=(0, 1, 3, 2))
# --- Encoder Path with Interlaced Layers ---
# We create all stages and downsamplers as ModuleLists
self.enc_stages = nn.ModuleList()
# This list defines the channel progression
# e.g., for ch=48: (48, 96, 144, 192, 240, 288)
channels = [ch * (i + 1) for i in range(num_stages + 1)]
for i in range(num_stages):
# Append the TDF_Block stage
self.enc_stages.append(TDF_Block(channels[i], dim_f // (2**i), freq_hidden_dim // (2**i), BN_EPS))
# Append the downsampling block immediately after
self.enc_stages.append(nn.Sequential(nn.Conv2d(channels[i], channels[i+1], 2, 2, bias=True), nn.ReLU(True)))
# --- Bottleneck ---
bottleneck_in_ch = channels[num_stages]
self.bottleneck = TDF_Block(bottleneck_in_ch, dim_f // (2**num_stages), freq_hidden_dim // (2**num_stages), BN_EPS)
# --- Decoder Path with Interlaced Layers ---
self.dec_stages = nn.ModuleList()
for i in range(num_stages):
dec_idx = num_stages - 1 - i
# Upsampler takes bottleneck/previous stage channels and outputs encoder stage channels
in_ch = channels[dec_idx + 1]
out_ch = channels[dec_idx]
# Append the upsampling block
seq = nn.Sequential(nn.ConvTranspose2d(in_ch, out_ch, 2, 2, bias=True),
nn.BatchNorm2d(out_ch, eps=BN_EPS),
nn.ReLU(True))
self.dec_stages.append(seq)
# Append the TDF_Block stage immediately after
self.dec_stages.append(TDF_Block(out_ch, dim_f // (2**dec_idx), freq_hidden_dim // (2**dec_idx), BN_EPS))
# --- Final Chain (always exists) ---
self.final_transpose = Transpose(dims=(0, 1, 3, 2))
self.final_conv = nn.Conv2d(ch, 4, 1, bias=True)
def forward(self, x):
# Initial processing
x = self.initial_conv(x)
x = self.initial_relu(x)
x = self.initial_transpose(x)
# --- Dynamic Encoder Path ---
skip_connections = []
# Encoder runs through the interlaced list
for i in range(0, self.num_stages*2, 2):
s = self.enc_stages[i](x)
skip_connections.append(s)
x = self.enc_stages[i+1](s)
# --- Bottleneck ---
x = self.bottleneck(x)
# --- Dynamic Decoder Path ---
skip_connections.reverse() # Reverse for easy lookup
# Decoder also runs through its interlaced list
for i in range(0, self.num_stages * 2, 2):
x = self.dec_stages[i](x)
x = x * skip_connections[i//2]
x = self.dec_stages[i+1](x)
# Final processing
output = self.final_transpose(x)
output = self.final_conv(output)
return output
View File
+104
View File
@@ -0,0 +1,104 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Wrappers for the model and inference
import logging
import torch
# ComfyUI imports
try:
import comfy.utils
with_comfy = True
except Exception:
with_comfy = False
# Local imports
from .stft import stft_chunk_process, stft_get_chunks
from ..db.load_model import load_model
from ..utils.misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.demixer")
SAMPLE_RATE = 44100
def show_inference_parameters(d):
logger.debug("Using inference parameters:")
logger.debug(f" Frequency Bins (n_fft/2): {d['mdx_n_fft_scale_set']//2}")
logger.debug(f" Amplitude Compensation: {d['compensate']}")
class DemixerMDX(object):
def __init__(self, d, device, models_dir):
self.d = d
self.model_run = load_model(d, device, models_dir)
self.device = device
show_inference_parameters(d)
self.sr = SAMPLE_RATE
self.ch = 2
def __call__(self, waveform, segments=1):
dim_t = (2 ** self.d['mdx_dim_t_set']) * segments
try:
# --- 1. Normalize input shape to handle both batched and non-batched data ---
if waveform.ndim == 2:
# Input is [C, samples], add a batch dimension to make it [1, C, samples]
logger.debug("Input is not batched. Adding a temporary batch dimension.")
waveform = waveform.unsqueeze(0)
input_was_batched = False
elif waveform.ndim == 3:
# Input is already batched [B, C, samples]
input_was_batched = True
else:
raise ValueError(f"Unsupported waveform shape: {waveform.shape}. Expected 2 or 3 dimensions.")
batch_size = waveform.shape[0]
logger.info("🎛️ Performing demix...")
# Lists to store the separated stems from each item in the batch
list_of_main_stems = []
list_of_complement_stems = []
# ComfyUI progress bar
progress_bar_ui = None
if with_comfy:
chunks = stft_get_chunks(waveform.shape[2], self.d['mdx_n_fft_scale_set'], segment_size=dim_t)
chunks *= batch_size
progress_bar_ui = comfy.utils.ProgressBar(chunks)
# --- 2. Iterate through the batch ---
for i, single_waveform in enumerate(waveform):
# single_waveform has shape [C, samples]
logger.debug(f"Processing item {i+1}/{batch_size}...")
# Process this single waveform
main_wav = stft_chunk_process(single_waveform, self.d, self.model_run, self.device, segment_size=dim_t,
progress_bar_ui=progress_bar_ui)
complement_wav = single_waveform - main_wav
# Add the results to our lists
list_of_main_stems.append(main_wav)
list_of_complement_stems.append(complement_wav)
# --- 3. Stack the results into single batch tensors ---
# torch.stack creates a new dimension (the batch dimension) from a list of tensors
stacked_main_stems = torch.stack(list_of_main_stems, dim=0)
stacked_complement_stems = torch.stack(list_of_complement_stems, dim=0)
# Both will now have shape [B, C, samples]
# --- 4. Denormalize output shape if original input was not batched ---
if not input_was_batched:
logger.debug("Squeezing batch dimension from output to match non-batched input.")
stacked_main_stems = stacked_main_stems.squeeze(0)
stacked_complement_stems = stacked_complement_stems.squeeze(0)
return [{'waveform': stacked_main_stems, 'sample_rate': SAMPLE_RATE, 'stem': self.d['primary_stem']},
{'waveform': stacked_complement_stems, 'sample_rate': SAMPLE_RATE, 'stem': 'Complement'}]
except Exception as e:
logger.error(f"Error during separation: {str(e)}")
raise e
def get_demixer(d, device, models_dir):
# Currently just MDX
return DemixerMDX(d, device, models_dir)
+21
View File
@@ -0,0 +1,21 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Helper to get a model from the correct class
import logging
from .MDX_Net import MDX_Net
from ..utils.misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.get_model")
# Currently we have just one type of networks, but this is a clean way to support more, or even test replacements
def get_model(d):
model_t = d['model_t'].lower()
if model_t != "mdx":
msg = f"Unknown model type `{model_t}`"
logger.error(msg)
raise ValueError(msg)
return MDX_Net(dim_f=d['mdx_dim_f_set'], ch=d['channels'], num_stages=d['stages'])
+164
View File
@@ -0,0 +1,164 @@
# Short-Time Fourier Transform (STFT).
import logging
import numpy as np
import torch
from tqdm import tqdm
# Local imports
from ..utils.misc import NODES_NAME
from ..utils.torch import model_to_target
logger = logging.getLogger(f"{NODES_NAME}.stft")
class STFT:
def __init__(self, n_fft, hop_length, dim_f, device):
self.n_fft = n_fft
self.hop_length = hop_length
self.window = torch.hann_window(window_length=self.n_fft, periodic=True)
self.dim_f = dim_f
self.device = device
def __call__(self, x):
window = self.window.to(x.device)
batch_dims = x.shape[:-2]
c, t = x.shape[-2:]
x = x.reshape([-1, t])
x = torch.stft(x, n_fft=self.n_fft, hop_length=self.hop_length,
window=window, center=True, return_complex=True)
x = torch.view_as_real(x)
x = x.permute([0, 3, 1, 2])
x = x.reshape([*batch_dims, c, 2, -1, x.shape[-1]]
).reshape([*batch_dims, c * 2, -1, x.shape[-1]])
return x[..., :self.dim_f, :]
# Original code
# def inverse(self, x):
# window = self.window.to(x.device)
# batch_dims = x.shape[:-3]
# c, f, t = x.shape[-3:]
# n = self.n_fft // 2 + 1
# f_pad = torch.zeros([*batch_dims, c, n - f, t]).to(x.device)
# x = torch.cat([x, f_pad], -2)
# x = x.reshape([*batch_dims, c // 2, 2, n, t]).reshape([-1, 2, n, t])
# x = x.permute([0, 2, 3, 1])
# x = x[..., 0] + x[..., 1] * 1.j
# x = torch.istft(x, n_fft=self.n_fft,
# hop_length=self.hop_length, window=window, center=True)
# x = x.reshape([*batch_dims, 2, -1])
#
# return x
# Annotated code
def inverse(self, x):
"""
Correctly performs the inverse STFT.
x is the output of the model, shape (B, C, F, T)
With C == 4 (L/R as complex)
"""
window = self.window.to(x.device)
batch_dims = x.shape[:-3]
c, f, t = x.shape[-3:] # c is 4 here
assert c == 4
n = self.n_fft // 2 + 1 # Full number of frequency bins
# Pad the frequency dimension back to its original size
f_pad = torch.zeros([*batch_dims, c, n - f, t], device=x.device)
x = torch.cat([x, f_pad], -2)
# The key is to correctly un-stack the 4 channels back into (C, 2)
# where C=2 (stereo) and 2 is real/imag.
# Reshape (B, 4, F, T) -> (B, 2, 2, F, T)
# The new dimensions are (B, stereo_channels, real_imag, F, T)
x = x.reshape([*batch_dims, 2, 2, n, t])
# Permute to get (B, stereo_channels, F, T, real_imag)
x = x.permute(0, 1, 3, 4, 2)
# Ensure the tensor is contiguous in memory before the final view
x = x.contiguous()
# Now, view_as_complex will work on the last dimension
x = torch.view_as_complex(x) # Shape: (B, C, F, T) complex
# Reshape for istft: (B, C, F, T) -> (B*C, F, T)
x = x.reshape(-1, n, t)
# Perform inverse STFT
x = torch.istft(x, n_fft=self.n_fft, hop_length=self.hop_length, window=window, center=True)
# Reshape back to (B, C, num_samples)
x = x.reshape([*batch_dims, 2, -1])
return x
def stft_get_chunks(samples, n_fft, segment_size=256, hop_length=1024):
chunk_size = hop_length * (segment_size - 1)
step = chunk_size - n_fft
return 1 + (samples - chunk_size + step - 1) // step
def stft_chunk_process(waveform, d, model_run, device, segment_size=256, hop_length=1024, progress_bar_ui=None):
""" Inference helper for models using STFT information as input """
n_fft = d['mdx_n_fft_scale_set']
compensate = d['compensate']
mix_np = waveform.numpy()
stft = STFT(n_fft, hop_length, model_run.dim_f, device)
trim = n_fft // 2
chunk_size = hop_length * (segment_size - 1)
gen_size = chunk_size - 2 * trim
# --- Overlap-Add Loop ---
pad = gen_size + trim - (mix_np.shape[1] % gen_size)
# Padded mixture as a numpy array
# mixture = np.concatenate((np.zeros((2, trim)), mix_np, np.zeros((2, pad))), axis=1)
mixture = np.concatenate((np.zeros((2, trim), dtype=np.float32), mix_np, np.zeros((2, pad), dtype=np.float32)), axis=1)
step = chunk_size - n_fft # Correct step size for large overlap
result = np.zeros((1, 2, mixture.shape[1]), dtype=np.float32)
divider = np.zeros((1, 2, mixture.shape[1]), dtype=np.float32)
total_chunks = 1 + (mixture.shape[1] - chunk_size + step - 1) // step
logger.info(f"⚙️ Processing {total_chunks} chunks...")
model_run.target_device = device
with model_to_target(model_run):
for i in tqdm(range(0, mixture.shape[1] - chunk_size + 1, step)):
start = i
end = i + chunk_size
mix_part = mixture[:, start:end]
# Convert just the chunk to a tensor for the model
mix_part_tensor = torch.from_numpy(mix_part).unsqueeze(0).to(device)
spek = stft(mix_part_tensor)
spec_pred = model_run(spek)
# Get the output back as a numpy array
tar_waves_np = stft.inverse(spec_pred).cpu().detach().numpy()
# Hanning window applied to the output before adding
# window = np.hanning(chunk_size)
window = np.hanning(chunk_size).astype(np.float32)
window = np.tile(window[None, None, :], (1, 2, 1))
result[..., start:end] += tar_waves_np * window
divider[..., start:end] += window
if progress_bar_ui:
progress_bar_ui.update(1)
# --- Final Normalization and Trimming ---
divider[divider == 0] = 1.0
main_wav_np = (result[0] / divider[0]) # Get the 2D array
main_wav_np = main_wav_np[:, trim:-trim][:, :mix_np.shape[1]]
main_wav_np *= compensate
# Convert final result back to a torch tensor for saving
return torch.from_numpy(main_wav_np)