diff --git a/models/uvr_model_data.json b/models/uvr_model_data.json index 6b2fd8f..555d64a 100644 --- a/models/uvr_model_data.json +++ b/models/uvr_model_data.json @@ -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", diff --git a/source/inference/HDemucs.py b/source/inference/HDemucs.py index b2e1a5e..eb87382 100644 --- a/source/inference/HDemucs.py +++ b/source/inference/HDemucs.py @@ -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: diff --git a/source/inference/demixer.py b/source/inference/demixer.py index c42ca34..47d639c 100644 --- a/source/inference/demixer.py +++ b/source/inference/demixer.py @@ -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. diff --git a/tool/demucs2safetensors.py b/tool/demucs2safetensors.py index 15b17b6..98f00de 100755 --- a/tool/demucs2safetensors.py +++ b/tool/demucs2safetensors.py @@ -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: