[Demucs][Added] Torchaudio model support

The one you get using HDEMUCS_HIGH_MUSDB_PLUS
Is an HDemucs model, the code in TorchAudio looks old, perhaps
simplified with incompatible renames.
This commit is contained in:
Salvador E. Tropea
2025-07-12 13:49:32 -03:00
parent c87923faeb
commit ab3b9d1c3d
4 changed files with 180 additions and 19 deletions
+30
View File
@@ -543,6 +543,21 @@
"primary_stem": "Vocals",
"stages": 5
},
"5ee7b421b6228a90e63855b108ad830b": {
"desc": "Hybrid Demucs MusDB High Train",
"download": "Main/Demucs",
"file_t": "safetensors",
"model_t": "Demucs",
"name": "hdemucs_high_trained.safetensors",
"params": "83639368",
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"use_demucs_pt_process": true
},
"5f6483271e1efb9bfb59e4a3e6d4d098": {
"channels": 32,
"compensate": 1.035,
@@ -660,6 +675,21 @@
"is_roformer": true,
"model_type": "SCNet"
},
"6f82220b9fc8445331d3a246f4ab509a": {
"desc": "Hybrid Demucs MusDB High Train",
"download": "Torch/Demucs",
"file_t": "pytorch",
"model_t": "Demucs",
"name": "hdemucs_high_trained.yaml",
"params": 83639368,
"primary_stem": [
"Vocals",
"Drums",
"Bass",
"Other"
],
"signatures": "[\"hdemucs_high_trained\"]"
},
"7205f4e53ec3226681ecd0c38e2fa747": {
"desc": "Hybrid Transformer Demucs",
"download": "Main/Demucs",
+7 -2
View File
@@ -70,7 +70,7 @@ class ScaledEmbedding(nn.Module):
class HEncLayer(nn.Module):
def __init__(self, chin, chout, kernel_size=8, stride=4, norm_groups=1, empty=False,
freq=True, dconv=True, norm=True, context=0, dconv_kw={}, pad=True,
rewrite=True):
rewrite=True, force_norm_in_last=False):
"""Encoder layer. This used both by the time and the frequency branch.
Args:
@@ -109,6 +109,9 @@ class HEncLayer(nn.Module):
pad = [pad, 0]
klass = nn.Conv2d
self.conv = klass(chin, chout, kernel_size, stride, pad)
if force_norm_in_last:
# PyTorch Audio uses it on last (empty) layer
self.norm1 = norm_fn(chout)
if self.empty:
return
self.norm1 = norm_fn(chout)
@@ -257,7 +260,7 @@ class MultiWrap(nn.Module):
class HDecLayer(nn.Module):
def __init__(self, chin, chout, last=False, kernel_size=8, stride=4, norm_groups=1, empty=False,
freq=True, dconv=True, norm=True, context=1, dconv_kw={}, pad=True,
context_freq=True, rewrite=True):
context_freq=True, rewrite=True, force_norm_in_last=False):
"""
Same as HEncLayer but for decoder. See `HEncLayer` for documentation.
"""
@@ -397,6 +400,7 @@ class HDemucs(nn.Module):
# Normalization
norm_starts=4,
norm_groups=4,
force_norm_in_last=False,
# DConv residual branch
dconv_mode=1,
dconv_depth=2,
@@ -533,6 +537,7 @@ class HDemucs(nn.Module):
kwt['kernel_size'] = kernel_size
kwt['stride'] = stride
kwt['pad'] = True
kwt['force_norm_in_last'] = force_norm_in_last
kw_dec = dict(kw)
multi = False
if multi_freqs and index < multi_freqs_depth:
+101 -14
View File
@@ -7,6 +7,7 @@
import logging
import math
import torch
import tqdm
# ComfyUI imports
try:
import comfy.utils
@@ -20,6 +21,8 @@ from ..utils.misc import NODES_NAME
from ..utils.torch import model_to_target, get_offload_device
from .demucs_api import apply_model, BagOfModels
from torchaudio.transforms import Fade
logger = logging.getLogger(f"{NODES_NAME}.demixer")
SAMPLE_RATE = 44100
@@ -135,6 +138,62 @@ def get_steps_for_demucs(model, wav, segment, shifts, overlap):
return math.ceil(wav.shape[-1] / stride) * (shifts + 1)
def separate_sources(
model: torch.nn.Module,
mix: torch.Tensor,
sample_rate: int,
segment: float = 10.0,
overlap: float = 0.1,
device: torch.device = None,
chunk_fade_shape: str = "linear",
) -> torch.Tensor:
"""
From: https://pytorch.org/audio/stable/tutorials/hybrid_demucs_tutorial.html
Apply model to a given mixture. Use fade, and add segments together in order to add model segment by segment.
Args:
segment (int): segment length in seconds
device (torch.device, str, or None): if provided, device on which to
execute the computation, otherwise `mix.device` is assumed.
When `device` is different from `mix.device`, only local computations will
be on `device`, while the entire tracks will be stored on `mix.device`.
"""
batch, channels, length = mix.shape
chunk_len = int(sample_rate * segment * (1 + overlap))
start = 0
end = chunk_len
overlap_frames = overlap * sample_rate
fade = Fade(fade_in_len=0, fade_out_len=int(overlap_frames), fade_shape=chunk_fade_shape)
chunks = math.ceil((length - overlap_frames) / chunk_len)
# Progress bars
progress_bar_console = tqdm.tqdm(total=chunks)
if with_comfy:
comfy_progress_bar = comfy.utils.ProgressBar(chunks)
final = torch.zeros(batch, len(model.sources), channels, length, device=device)
while start < length - overlap_frames:
chunk = mix[:, :, start:end]
progress_bar_console.update(1)
if with_comfy:
comfy_progress_bar.update(1)
with torch.no_grad():
out = model.forward(chunk)
out = fade(out)
final[:, :, :, start:end] += out
if start == 0:
fade.fade_in_len = int(overlap_frames)
start += int(chunk_len - overlap_frames)
else:
start += chunk_len
end += chunk_len
if end >= length:
fade.fade_out_len = 0
return final
class DemixerDemucs(DemixerGeneric):
def __init__(self, d, device, models_dir):
super().__init__(d, device, models_dir)
@@ -194,21 +253,49 @@ class DemixerDemucs(DemixerGeneric):
forced_segment = None
else:
logger.debug(f"Using user provided segment size {forced_segment} s")
if with_comfy:
comfy_progress_bar = comfy.utils.ProgressBar(self.get_steps(waveform_tensor, forced_segment, shifts, overlap))
with model_to_target(model):
separated_tensors = apply_model(
model,
input_tensor_on_device,
device=self.device,
segment=forced_segment,
shifts=shifts + 1, # Shifts 0 is disabled, 1 is just one pass, 2 is 2 passes
overlap=overlap,
split=True, # Enable chunking
progress=True, # Show a progress bar in the console
callback=lambda x: (x.get('state') == 'end') and comfy_progress_bar.update(1) if with_comfy else None,
)
# Normalize the input tensors
ref = input_tensor_on_device.mean(1, keepdim=True)
mean = ref.mean(dim=2, keepdim=True)
std = ref.std(dim=2, keepdim=True)
input_tensor_on_device = (input_tensor_on_device - mean) / (std + 1e-8)
if self.d.get('use_demucs_pt_process', False):
# This is for the model from torchaudio, using the Demucs code we get some strange noises
# Using their example they aren't produced
model.target_device = self.device
logger.debug("Using PyTorch Audio chunking for old model")
with model_to_target(model):
separated_tensors = separate_sources(
model,
input_tensor_on_device,
model.samplerate,
segment=forced_segment or 16.0,
overlap=overlap,
device=self.device,
chunk_fade_shape="half_sine"
)
else:
if with_comfy:
comfy_progress_bar = comfy.utils.ProgressBar(self.get_steps(waveform_tensor, forced_segment, shifts,
overlap))
with model_to_target(model):
separated_tensors = apply_model(
model,
input_tensor_on_device,
device=self.device,
segment=forced_segment,
shifts=shifts + 1, # Shifts 0 is disabled, 1 is just one pass, 2 is 2 passes
overlap=overlap,
split=True, # Enable chunking
progress=True, # Show a progress bar in the console
callback=lambda x: (x.get('state') == 'end') and comfy_progress_bar.update(1) if with_comfy else None,
)
# Denormalize the output stems
mean = mean.unsqueeze(1)
std = std.unsqueeze(1)
separated_tensors = separated_tensors * std + mean
# Move the final result tensor back to the CPU before creating the output dicts.
# This is good practice to free up VRAM for subsequent nodes.
+42 -3
View File
@@ -26,7 +26,7 @@ try:
except Exception:
with_demuc_lib = False
import bootstrap # noqa: F401
from source.utils.misc import cli_add_verbose, FractionEncoder
from source.utils.misc import cli_add_verbose, FractionEncoder, debugl
from source.utils.logger import main_logger, logger_set_standalone
from source.db.models_db import cli_add_db, get_download_url, ModelsDB
from source.db.hash import get_hash
@@ -37,6 +37,10 @@ MODULES_MAP = {'demucs': local_demucs_module,
'demucs.demucs': local_demucs_module,
'demucs.hdemucs': local_hdemucs_module,
'demucs.htdemucs': local_htdemucs_module}
MAP = {'freq_encoder': 'encoder',
'freq_decoder': 'decoder',
'time_encoder': 'tencoder',
'time_decoder': 'tdecoder'}
logger = main_logger
@@ -62,6 +66,36 @@ def remap_module(modules_map):
del sys.modules[old_name]
def solve_simple_pt(yaml_data, pkg):
""" PyTorch Audio lib has some raw models, we store the metadata in the YAML """
klass = yaml_data.get('klass')
if klass is None:
main_logger.error("No `klass` in YAML")
sys.exit(4)
if klass == 'Demucs':
klass = local_demucs_module.Demucs
elif klass == 'HDemucs':
klass = local_hdemucs_module.HDemucs
elif klass == 'HTDemucs':
klass = local_htdemucs_module.HTDemucs
else:
main_logger.error("Unknown model `klass` {klass}")
sys.exit(4)
args = yaml_data.get('args', {})
kwargs = yaml_data.get('kwargs', {})
# For PyTorch Audio model (very old code?)
new_dict = {}
for k, v in pkg.items():
parts = k.split('.')
gr = parts[0]
if gr in MAP:
k = MAP[gr] + '.' + '.'.join(parts[1:])
new_dict[k] = v
return klass, args, kwargs, new_dict
def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path_str: str, all_metadata, data: Dict):
"""
Loads an original Demucs model bag, extracts weights and all necessary metadata
@@ -112,7 +146,11 @@ def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path
# This is the only "unsafe" part, loading the original pickle file
pkg = torch.load(path, map_location='cpu', weights_only=False)
klass, args, kwargs, state = pkg["klass"], pkg["args"], pkg["kwargs"], pkg["state"]
debugl(logger, 2, f"PyTorch data type is {type(pkg)}")
if isinstance(pkg, dict) and 'klass' in pkg:
klass, args, kwargs, state = pkg["klass"], pkg["args"], pkg["kwargs"], pkg["state"]
else:
klass, args, kwargs, state = solve_simple_pt(yaml_data, pkg)
main_logger.info(f" - Model class: {klass.__module__}.{klass.__name__}")
# Dequantize if necessary by letting the original code handle it
@@ -163,6 +201,7 @@ def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path
if __name__ == '__main__':
parser = argparse.ArgumentParser(description="Convert Demucs .th models to a single .safetensors file.")
parser.add_argument('--yaml', required=True, type=str, help="Path to the Demucs .yaml file.")
parser.add_argument('--models', nargs='*', default=None, type=str,
help="Optional. Paths to .th files. If not provided, assumes they are "
@@ -206,7 +245,7 @@ if __name__ == '__main__':
model_files = []
for sig in set(signatures): # Use set to avoid redundant searches
found = list(yaml_dir.glob(f'{sig}*.th'))
found = list(yaml_dir.glob(f'{sig}*.th')) + list(yaml_dir.glob(f'{sig}*.pt'))
if not found:
raise FileNotFoundError(f"Could not automatically find a model file for signature '{sig}' in {yaml_dir}")
if len(found) > 1: