[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:
@@ -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",
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user