diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..40a1ef7 --- /dev/null +++ b/__init__.py @@ -0,0 +1,35 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +import inspect +import logging +from .source.utils.misc import NODES_NAME +from . import nodes # noqa: E402 + +init_logger = logging.getLogger(f"{NODES_NAME}.__init__") + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + + +def register_nodes(module): + suffix = " " + module.SUFFIX if hasattr(module, "SUFFIX") else "" + if suffix: + suffix = " " + suffix + for name, obj in inspect.getmembers(module): + if not inspect.isclass(obj) or not hasattr(obj, "INPUT_TYPES"): + continue + assert hasattr(obj, "UNIQUE_NAME"), f"No name for {obj.__name__}" + NODE_CLASS_MAPPINGS[obj.UNIQUE_NAME] = obj + NODE_DISPLAY_NAME_MAPPINGS[obj.UNIQUE_NAME] = obj.DISPLAY_NAME + suffix + + +register_nodes(nodes) + +init_logger.info(f"Registering {len(NODE_CLASS_MAPPINGS)} node(s).") +init_logger.debug(f"{list(NODE_DISPLAY_NAME_MAPPINGS.values())}") + +WEB_DIRECTORY = "./js" +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..e8d63df --- /dev/null +++ b/nodes.py @@ -0,0 +1,143 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +import os +import torch +from typing import Dict +# ComfyUI imports +import folder_paths # ComfyUI's way to access model paths +# Local imports +from .source.utils.logger import main_logger +from .source.utils.load_audio import audio_get_channels, force_stereo, force_sample_rate +from .source.utils.torch import get_torch_device_options +from .source.utils.comfy_node_action import send_node_action +from .source.inference.demixer import get_demixer +from .source.db.models_db import ModelsDB + + +DEF_MODEL = 'Kim_Vocal_2.safetensors' +DEF_ENTRY = 'Default' +MODELS_DIR = os.path.join(folder_paths.models_dir, "audio", "MDX") +models_db = ModelsDB(MODELS_DIR) +logger = main_logger + + +class AudioSeparateVocals: + PRIMARY_STEM = 'Vocals' + MODEL_T = 'MDX' + FILE_T = 'safetensors' + DEFAULT_MODEL = "Kim_Vocal_2.safetensors" + + @classmethod + def _get_available_audio_models(cls): + global models_db + # Refresh the database + models_db.refresh() + # Filter the models this node can handle + cls.models_filtered = models_db.get_filtered(primary_stem=cls.PRIMARY_STEM, model_t=cls.MODEL_T, file_t=cls.FILE_T, + default=cls.DEFAULT_MODEL, repeat_dl=True) + # We add any model downloaded and memorized by the GUI + return cls.models_filtered.get_display_names() + + @classmethod + def INPUT_TYPES(cls): + device_options, default_device = get_torch_device_options() + return { + "required": { + "input_sound": ("AUDIO",), + "model": (cls._get_available_audio_models(),), # Dropdown for model selection + "segments": ("INT", { + "default": 1, # Default value + "min": 1, # Minimum allowed value + "max": 64, # Maximum allowed value (set a reasonable practical max) + "step": 1, # Step for slider/spinbox + "display": "slider" # How to display: "number" or "slider" + }), + "target_device": (device_options, { + "default": default_device, + "tooltip": "The device (CPU or CUDA) to which the projection layer will be assigned for computation."}), + } + } + + RETURN_TYPES = ("AUDIO", "AUDIO",) + RETURN_NAMES = (PRIMARY_STEM, "Complement",) + FUNCTION = "execute" + CATEGORY = "audio/separation" + DESCRIPTION = "Separates vocals using MDX-Net networks" + UNIQUE_NAME = "AudioSeparateVocals" + DISPLAY_NAME = "Vocals using MDX" + + def __init__(self): + super().__init__() + self.demixer = None + + def execute(self, input_sound: Dict, model: str, segments: int, target_device: str): + # Get information for the selected model + main_logger.info(f"Selected model: {model}") + model_data = self.models_filtered.get_by_display_name(model) + if model_data is None: + raise ValueError("Unknown model selected, please refresh pressing `R` and select another") + model_path = model_data.get('model_path') + + # Create or recycle a demixer + device = torch.device(target_device) + if self.demixer is None or self.demixer.d['hash'] != model_data['hash']: + # New demixer + logger.debug("Creating a new demixer object") + # This will load the model, optionally downloading it + self.demixer = get_demixer(model_data, device, MODELS_DIR) + + # Handle a change in the icon of the model name + if model_path is None: + # Was downloaded + send_node_action("change_widget", "model", model_data['indicator'] + model_data['filtered_name']) + + # Match channels and S/R + waveform = input_sound['waveform'] + sample_rate = input_sound['sample_rate'] + if audio_get_channels(waveform) == 1 and self.demixer.ch == 2: + waveform = force_stereo(waveform) + if sample_rate != self.demixer.sr: + waveform = force_sample_rate(waveform, sample_rate, self.demixer.sr) + + # Demix + wavs = self.demixer(waveform, segments) + + return (wavs[0], wavs[1],) + + +class AudioSeparateInstrumental(AudioSeparateVocals): + PRIMARY_STEM = 'Instrumental' + DEFAULT_MODEL = "Kim_Inst.safetensors" + DESCRIPTION = "Separates instruments using MDX-Net networks" + UNIQUE_NAME = "AudioSeparateInstrumental" + DISPLAY_NAME = "Instrumental using MDX" + RETURN_NAMES = (PRIMARY_STEM, "Complement",) + + +class AudioSeparateBass(AudioSeparateVocals): + PRIMARY_STEM = 'Bass' + DEFAULT_MODEL = "kuielab_b_bass.safetensors" + DESCRIPTION = "Separates bass using MDX-Net networks" + UNIQUE_NAME = "AudioSeparateBass" + DISPLAY_NAME = "Bass using MDX" + RETURN_NAMES = (PRIMARY_STEM, "Complement",) + + +class AudioSeparateDrums(AudioSeparateVocals): + PRIMARY_STEM = 'Drums' + DEFAULT_MODEL = "kuielab_b_drums.safetensors" + DESCRIPTION = "Separates drums using MDX-Net networks" + UNIQUE_NAME = "AudioSeparateDrums" + DISPLAY_NAME = "Drums using MDX" + RETURN_NAMES = (PRIMARY_STEM, "Complement",) + + +class AudioSeparateVarious(AudioSeparateVocals): + PRIMARY_STEM = ["Other", "Reverb"] + DEFAULT_MODEL = "Reverb_HQ_By_FoxJoy.safetensors" + DESCRIPTION = "Misc. separators using MDX-Net networks" + UNIQUE_NAME = "AudioSeparateVarious" + DISPLAY_NAME = "Various using MDX" + RETURN_NAMES = ("Main", "Complement",) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7bc38ba --- /dev/null +++ b/requirements.txt @@ -0,0 +1,8 @@ +torch +torchaudio +numpy +tqdm +# Optionals: +# requests +# colorama +# onnxruntime diff --git a/source/__init__.py b/source/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/source/db/__init__.py b/source/db/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/source/db/hash.py b/source/db/hash.py new file mode 100644 index 0000000..0fcc35d --- /dev/null +++ b/source/db/hash.py @@ -0,0 +1,48 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# From various UVR clones adapted by Gemini 2.5 Pro +import hashlib +import re + +HASH_REGEX = re.compile(r'^[\da-z]{32}$') +# The size to seek from the end of the file, in bytes. +# 10000 * 1024 bytes = 10,000 KB = ~9.77 MB +SEEK_SIZE = 10000 * 1024 + + +def get_hash(filepath): + """ + Calculates the MD5 hash of a file using a special method. + + It tries to hash only the last SEEK_SIZE bytes of the file. This is a + shortcut used by some communities (e.g., for large AI models) to get a + quick, unique hash without processing the entire file. + + If the file is smaller than SEEK_SIZE, it falls back to hashing the + entire file. + + Args: + filepath (str): The path to the file. + + Returns: + str: The calculated MD5 hexdigest. + """ + try: + with open(filepath, 'rb') as f: + # Seek to SEEK_SIZE bytes from the end of the file (whence=2) + f.seek(-SEEK_SIZE, 2) + file_hash = hashlib.md5(f.read()).hexdigest() + except (IOError, OSError): + # This will happen if the file is smaller than SEEK_SIZE. + # In that case, hash the entire file. + with open(filepath, 'rb') as f: + file_hash = hashlib.md5(f.read()).hexdigest() + + return file_hash + + +def is_hash(value): + return HASH_REGEX.match(value) diff --git a/source/db/hash_dir.py b/source/db/hash_dir.py new file mode 100644 index 0000000..0630792 --- /dev/null +++ b/source/db/hash_dir.py @@ -0,0 +1,227 @@ +#!/usr/bin/env python3 +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Tool to get hashes for all the files in a dir, using a cache +# Can be used as an standalone tool for testing +# Code by Gemini 2.5 Pro +import os +import csv +import logging +import sys + +from .hash import get_hash +from ..utils.misc import NODES_NAME, debugl + +# Set up the logger as specified +logger = logging.getLogger(f"{NODES_NAME}.hash_dir") + +# Constants for clarity +MIN_FILE_SIZE_MB = 10 +MIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_MB * 1024 * 1024 +CATALOG_FILENAME = ".catalog.csv" + + +def hash_dir(directory_path: str) -> dict: + """ + Computes hashes for large files in a directory, using a local cache + to avoid re-computation. + + Args: + directory_path: The path to the directory to scan. + + Returns: + A dictionary mapping {file_hash: file_name} for all files in the + directory that are 10MB or larger. + """ + logger.debug(f"Starting hash process for directory: '{directory_path}'") + + # 1. Ensure the target directory exists. + try: + os.makedirs(directory_path, exist_ok=True) + except OSError as e: + logger.error(f"FATAL: Could not create directory '{directory_path}'. Error: {e}") + # This is a fatal error, so we raise the exception to stop execution. + raise + + catalog_path = os.path.join(directory_path, CATALOG_FILENAME) + cached_data = {} + + # 2. Load existing cache from ".catalog.csv" if it exists. + try: + with open(catalog_path, 'r', newline='') as f: + reader = csv.reader(f) + # Skip header + next(reader, None) + for row in reader: + if len(row) == 3: + filename, file_hash, timestamp = row + cached_data[filename] = (file_hash, float(timestamp)) + logger.debug(f"Successfully loaded {len(cached_data)} entries from cache: {catalog_path}") + except FileNotFoundError: + logger.debug(f"Cache file '{catalog_path}' not found. A new one will be created.") + except Exception as e: + logger.warning(f"Could not read cache file '{catalog_path}'. Proceeding without cache. Error: {e}") + + final_hashes = {} + updated_cache = {} + has_updates = False + + # 3. Iterate through files in the directory. + for filename in os.listdir(directory_path): + file_path = os.path.join(directory_path, filename) + + # Skip subdirectories and the cache file itself + if not os.path.isfile(file_path) or filename == CATALOG_FILENAME: + continue + + # 4. Filter by file size. + try: + file_size = os.path.getsize(file_path) + if file_size < MIN_FILE_SIZE_BYTES: + logger.debug(f"Skipping small file: '{filename}' ({file_size / 1024**2:.2f}MB)") + continue + except OSError as e: + logger.warning(f"Could not get size of file '{filename}'. Skipping. Error: {e}") + continue + + current_mtime = os.path.getmtime(file_path) + file_hash = None + + # 5. Check cache for a valid, up-to-date entry. + if filename in cached_data: + cached_hash, cached_mtime = cached_data[filename] + if current_mtime == cached_mtime: + debugl(logger, 2, f"Cache hit for '{filename}'. Using stored hash.") + file_hash = cached_hash + else: + logger.debug(f"File '{filename}' has been modified. Re-calculating hash.") + + # 6. If no valid cache entry, compute the hash. + if file_hash is None: + logger.debug(f"Computing hash for '{filename}'...") + try: + file_hash = get_hash(file_path) + has_updates = True + except IOError as e: + logger.warning(f"Could not read file '{filename}' to compute hash. Skipping. Error: {e}") + continue + + # Add to the results and prepare for caching + final_hashes[file_hash] = file_path + updated_cache[filename] = (file_hash, current_mtime) + + # 7. Write the updated cache back to disk. + if has_updates: + try: + with open(catalog_path, 'w', newline='') as f: + writer = csv.writer(f) + writer.writerow(['Filename', 'Hash', 'Time stamp']) + for filename, (file_hash, mtime) in updated_cache.items(): + writer.writerow([filename, file_hash, mtime]) + logger.debug(f"Successfully wrote {len(updated_cache)} entries to cache file '{catalog_path}'.") + except IOError as e: + # This is a non-fatal warning as per requirements. + logger.warning(f"Failed to write to cache file '{catalog_path}'. Hashes were computed but not saved. Error: {e}") + + logger.debug(f"Hash process finished. Found {len(final_hashes)} valid files.") + return final_hashes + + +# ============================================================================== +# Command-Line Tool for Testing and Operation +# python -m source.db.hash_dir +# ============================================================================== +if __name__ == "__main__": + # Local imports to avoid top-level pollution + import argparse + from pathlib import Path + import pprint + + # --- Argument Parsing --- + parser = argparse.ArgumentParser( + description="Calculate and cache hashes for large files in a directory.", + formatter_class=argparse.RawTextHelpFormatter + ) + parser.add_argument( + 'directory', + nargs='?', # Makes the argument optional; we'll validate it manually. + help='The directory to scan. Required unless --test is used.' + ) + parser.add_argument( + '--test', + action='store_true', # This is a flag; if present, args.test will be True. + help='Run in test mode. Creates a temporary directory with dummy files and runs a test sequence.' + ) + args = parser.parse_args() + + # --- Setup Logging --- + # Enable debug logging for command-line execution to provide useful feedback. + logging.basicConfig( + level=logging.DEBUG, + format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' + ) + + # --- Main Logic --- + try: + if args.test: + # --- TEST MODE --- + # This mode ignores the 'directory' argument and uses a fixed test path. + + # Helper function only needed for test mode + def create_test_file(path: Path, size_in_bytes: int): + """Helper to create a dummy file of a specific size.""" + path.parent.mkdir(parents=True, exist_ok=True) + with open(path, "wb") as f: + f.write(os.urandom(size_in_bytes)) + logger.info(f"Created test file: {path} ({size_in_bytes / 1024**2:.2f}MB)") + + test_dir = Path("./temp_hash_dir_test_cli") + print("-" * 60) + print("Running in TEST mode.") + print(f"Test directory: '{test_dir.resolve()}'") + print("-" * 60) + + # Create test files + create_test_file(test_dir / "large_file_1.bin", 11 * 1024 * 1024) # > 10MB + create_test_file(test_dir / "large_file_2.bin", 15 * 1024 * 1024) # > 10MB + create_test_file(test_dir / "small_file.txt", 1 * 1024 * 1024) # < 10MB + create_test_file(test_dir / "edge_case_file.bin", 10 * 1024 * 1024) # exactly 10MB + + # Run the test sequence + print("\n>>> FIRST RUN: Computing all hashes...") + result_dict = hash_dir(str(test_dir)) + pprint.pprint(result_dict) + + print("\n>>> SECOND RUN: Should use cache for all files...") + result_dict_cached = hash_dir(str(test_dir)) + pprint.pprint(result_dict_cached) + + print("\n>>> THIRD RUN: Modifying a file and re-running to test cache invalidation...") + create_test_file(test_dir / "large_file_1.bin", 12 * 1024 * 1024) + result_dict_modified = hash_dir(str(test_dir)) + pprint.pprint(result_dict_modified) + + else: + # --- NORMAL OPERATION MODE --- + + # In normal mode, the 'directory' argument is required. + if not args.directory: + parser.error("The 'directory' argument is required when not using --test.") + + target_directory = args.directory + print("-" * 60) + print(f"Running in NORMAL mode on directory: '{target_directory}'") + print("-" * 60) + + # Just run the function once on the specified directory + result_dict = hash_dir(target_directory) + + print("\n>>> Resulting Hashes:") + pprint.pprint(result_dict) + + except Exception as e: + logger.critical(f"A critical error occurred during execution: {e}") + sys.exit(1) diff --git a/source/db/load_model.py b/source/db/load_model.py new file mode 100644 index 0000000..9f2f9e0 --- /dev/null +++ b/source/db/load_model.py @@ -0,0 +1,52 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Model loader helper +# Original code from Gemini 2.5 Pro +import logging +# Local imports +from .models_db import download_model +from ..inference.get_model import get_model +from ..utils.misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.load_model") + + +def show_model_parameters(d): + logger.debug("Using model parameters:") + logger.debug(f" Frequency Dimension (dim_f): {d['mdx_dim_f_set']}") + logger.debug(f" Base Channels (ch): {d['channels']}") + logger.debug(f" U-Net Stages: {d['stages']}") + + +def load_model(d, device, models_dir): + file_t = d['file_t'].lower() + show_model_parameters(d) + + # Get the file name, download if necessary + model_path = d.get('model_path') + if model_path is None: + # It means it wasn't on disk + model_path = download_model(d, models_dir) + + # ONNX + if file_t == "onnx": + from ..utils.load_onnx import load_onnx + model = load_onnx(model_path, device) + # Store the same information we have in the PyTorch version + model.dim_f = d['mdx_dim_f_set'] + model.ch = d['channels'] + model.num_stages = d['stages'] + return model + + # Safetensors + if file_t == "safetensors": + from ..utils.load_safetensors import load_safetensors + return load_safetensors(model_path, get_model(d), device) + + # Other + msg = f"Unknown file type {file_t}" + logger.error(msg) + raise ValueError(msg) diff --git a/source/db/models_db.py b/source/db/models_db.py new file mode 100644 index 0000000..0b71df1 --- /dev/null +++ b/source/db/models_db.py @@ -0,0 +1,317 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Handles a JSON file with information about the supported models. +# The JSON file is similar to what UVR uses, but with more information. +import json +import logging +import os +from pathlib import Path +from ..utils.misc import NODES_NAME +from ..utils.downloader import download_model as download_model_basic +from ..utils.comfy_notification import send_toast_notification +from .hash_dir import hash_dir +from .hash import is_hash, get_hash + +logger = logging.getLogger(f"{NODES_NAME}.models_db") +known_models = None +known_models_mtime = None +# ICON_REMOTE = "\u2601" # ☁️ Cloud +ICON_REMOTE = "⬇️ " # "\u2B07" # ⬇️ +ICON_DOWNLOADED = "\U0001F4BE " # 💾 Floppy Disk +KNOWN_SOURCES = {'Politrees/MDXNet': 'https://huggingface.co/Politrees/UVR_resources/resolve/main/models/MDXNet', + 'Main/MDX': 'https://huggingface.co/set-soft/audio_separation/resolve/main/MDX'} + + +def get_db_filename(provided=None): + if provided is not None: + return provided + # Get the directory where this script is located + script_dir = Path(__file__).resolve().parent + + # Build the path to the JSON file: go up one level, then into models/ + json_path = script_dir / ".." / ".." / "models" / "uvr_model_data.json" + try: + return json_path.resolve().relative_to(Path.cwd()) + except ValueError: + pass + return json_path.resolve() + + +def load_known_models(json_path=None): + """ + Loads the uvr_model_data.json configuration file using a path relative + to this script's location. This is the most reliable method. + """ + global known_models + global known_models_mtime + json_path = get_db_filename(json_path) + + # Check if we have a fresh db + if known_models is not None and known_models_mtime == os.path.getmtime(json_path): + return known_models + + try: + logger.debug(f"Attempting to load JSON from: {json_path}") + + # Open and load the JSON file + with open(json_path, 'r', encoding='utf-8') as f: + data = json.load(f) + + known_models = data + known_models_mtime = os.path.getmtime(json_path) + return data + + except FileNotFoundError: + logger.error("Error: The models database was not found at the expected location.") + logger.error("Please check the directory structure.") + return None + except json.JSONDecodeError: + logger.error(f"Error: The file at '{json_path}' is not a valid JSON file.") + return None + except Exception as e: + logger.error(f"An unexpected error occurred: {e}") + return None + + +def save_known_models(data, json_path=None): + json_path = get_db_filename(json_path) + backup_path = os.path.join(os.path.dirname(json_path), f".{os.path.basename(json_path)}~") + + try: + logger.debug(f"Creating DB backup at '{backup_path}'...") + if os.path.exists(json_path): + os.rename(json_path, backup_path) + + logger.debug(f"Saving new database to '{json_path}'...") + with open(json_path, 'w') as f: + json.dump(data, f, indent=4, sort_keys=True) + + logger.debug("✅ Database updated and saved successfully.") + + global known_models + global known_models_mtime + known_models = data + known_models_mtime = os.path.getmtime(json_path) + + except Exception as e: + logger.error(f"during database save: {e}") + logger.info("Attempting to restore from backup...") + if os.path.exists(backup_path): + os.rename(backup_path, json_path) + logger.info("Backup restored.") + raise + + +def create_display_name(d, hash, desc, primary_stem, model_t, file_t, downloaded): + if primary_stem and len(primary_stem) == 1: + desc = desc.replace(next(iter(primary_stem)), "") + if model_t and len(model_t) == 1: + desc = desc.replace(next(iter(model_t)), "") + if file_t is None or len(file_t) > 1: + # We can get a name collision i.e. ONNX + safetensors + file_extra = f" [{d['file_t']}]" + else: + file_extra = "" + return " ".join(desc.split()) + file_extra + + +def get_models_full(primary_stem=None, model_t=None, file_t=None, json_path=None, downloaded=None, default=None, + repeat_dl=False): + """ Returns a dict with models that satisfy the provided criteria """ + + # Allow for multiple values in the filters + if isinstance(primary_stem, str): + primary_stem = {primary_stem} + if isinstance(model_t, str): + model_t = {model_t} + if isinstance(file_t, str): + file_t = {file_t} + + # Filter the models db + models = load_known_models(json_path) + found = {} + found_hashes = {} + found_disk = {} + def_sep = [] + on_disk = [] + to_down = [] + on_disk_as_down = [] + for hash, d in models.items(): + # Skip unnamed models, valid, these are models we don't have or support + try: + name = d['name'] + except KeyError: + continue + # Description is mandatory + try: + desc = d['desc'] + except KeyError: + logger.error(f"Missing `desc` for {name}") + continue + # Check the stem + try: + if primary_stem is not None and d['primary_stem'] not in primary_stem: + continue + except KeyError: + logger.error(f"Missing `primary_stem` for {name}") + continue + # Check the network type + try: + if model_t is not None and d['model_t'] not in model_t: + continue + except KeyError: + logger.error(f"Missing `model_t` for {name}") + continue + # Check the container type + try: + if file_t is not None and d['file_t'] not in file_t: + continue + except KeyError: + logger.error(f"Missing `file_t` for {name}") + continue + filtered_name = create_display_name(d, hash, desc, primary_stem, model_t, file_t, downloaded) + d['hash'] = hash + d['filtered_name'] = filtered_name + if downloaded: + file_name = downloaded.get(hash) + d['indicator'] = ICON_DOWNLOADED if file_name is not None else ICON_REMOTE + else: + file_name = None + d['indicator'] = "" + if default and default == name: + if downloaded is None: + def_sep.append(filtered_name) + else: + if file_name is not None: + def_sep.append(ICON_DOWNLOADED + filtered_name) + if repeat_dl: + # A copy for ComfyUI, so the user doesn't need to refresh the node and choose the one downloaded + # Also helps to allow forcing save a node with the file as "not downloaded" + # Is a hack, but is the best I came up + on_disk_as_down.append(ICON_REMOTE + filtered_name) + else: + def_sep.append(ICON_REMOTE + filtered_name) + else: + if file_name is not None: + on_disk.append(ICON_DOWNLOADED + filtered_name) + if repeat_dl: + on_disk_as_down.append(ICON_REMOTE + filtered_name) + else: + to_down.append(ICON_REMOTE + filtered_name) + found[filtered_name] = d + found_hashes[hash] = d + if file_name is not None: + d['model_path'] = downloaded[hash] + found_disk[os.path.realpath(file_name)] = d + + return found, found_hashes, found_disk, def_sep + sorted(on_disk) + sorted(to_down) + sorted(on_disk_as_down) + + +def get_models(primary_stem=None, model_t=None, file_t=None, json_path=None, downloaded=None, default=None, by_hash=False): + dnames, hashes, fnames, ldnames = get_models_full(primary_stem, model_t, file_t, json_path, downloaded, default) + return hashes if by_hash else dnames, ldnames + + +def get_download_url(data): + try: + name = data['name'] + dn_t = data['download'] + except KeyError: + return None + try: + return os.path.join(KNOWN_SOURCES[dn_t], name) + except KeyError: + logger.error(f"Unknown download source `{dn_t}`") + return None + + +def cli_add_models_and_db(parser): + # Compute the models dir assuming the script is run from a clone of the repo + default_json_file = get_db_filename() + + parser.add_argument('--models_dir', type=str, default=os.path.dirname(default_json_file), + help="Path to the directory containing model files.") + parser.add_argument('--json_file', type=str, default=default_json_file, + help="Path to the models database JSON file.") + + +def download_model(data, models_dir): + # Check we can download it + url = get_download_url(data) + if url is None: + raise ValueError("Model is not downloadable") + + # Download the file + name = data['name'] + send_toast_notification(f"Downloading `{name}`", "Download") + try: + fname = download_model_basic(url, models_dir, name) + # Mark it as downloaded + data['model_path'] = fname + if data['indicator']: + data['indicator'] = ICON_DOWNLOADED + # Notify the user + send_toast_notification("Finished downloading", "Download", 'success') + return fname + except Exception as e: + raise ValueError(f"Failed to download {name} from {url}\n{e}") + + +class FilteredModels(object): + def __init__(self, primary_stem=None, model_t=None, file_t=None, json_path=None, downloaded=None, default=None, + repeat_dl=False): + self.primary_stem = primary_stem + self.model_t = model_t + self.file_t = file_t + self.by_dname, self.by_hash, self.by_fname, self.dnames = get_models_full(primary_stem, model_t, file_t, json_path, + downloaded, default, repeat_dl) + + def get_by_display_name(self, name): + if name.startswith(ICON_REMOTE) or name.startswith(ICON_DOWNLOADED): + name = name[name.index(' ')+1:].strip() + return self.by_dname.get(name) + + def get_by_hash(self, hash): + return self.by_hash.get(hash) + + def get_by_file_name(self, file): + return self.by_fname.get(file) + + def get(self, value): + # If it looks like a hash try it first + if is_hash(value): + d = self.by_hash.get(value) + if d: + return d + # Then try by a display name + d = self.by_dname.get(value) + if d: + return d + # Is this a file name? + if os.path.isfile(value): + # Try using its hash + hash = get_hash(value) + return self.by_hash.get(hash) + return None + + def get_display_names(self, clean=False): + return [m[m.index(' ')+1:].strip() for m in self.dnames] if clean else self.dnames + + +class ModelsDB(object): + def __init__(self, models_dir: str, json_path: str = None): + super().__init__() + self.models_dir = models_dir + self.json_path = json_path + self.refresh() + + def refresh(self): + self.downloaded = hash_dir(self.models_dir) + self.models = load_known_models(self.json_path) + + def get_filtered(self, primary_stem=None, model_t=None, file_t=None, default=None, repeat_dl=False): + return FilteredModels(primary_stem=primary_stem, model_t=model_t, file_t=file_t, json_path=self.json_path, + downloaded=self.downloaded, default=default, repeat_dl=repeat_dl) diff --git a/source/inference/MDX_Net.py b/source/inference/MDX_Net.py new file mode 100644 index 0000000..61e7c44 --- /dev/null +++ b/source/inference/MDX_Net.py @@ -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 diff --git a/source/inference/__init__.py b/source/inference/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/source/inference/demixer.py b/source/inference/demixer.py new file mode 100644 index 0000000..93b1731 --- /dev/null +++ b/source/inference/demixer.py @@ -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) diff --git a/source/inference/get_model.py b/source/inference/get_model.py new file mode 100644 index 0000000..253b121 --- /dev/null +++ b/source/inference/get_model.py @@ -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']) diff --git a/source/inference/stft.py b/source/inference/stft.py new file mode 100644 index 0000000..0a18bea --- /dev/null +++ b/source/inference/stft.py @@ -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) diff --git a/source/utils/__init__.py b/source/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/source/utils/ansi.py b/source/utils/ansi.py new file mode 100644 index 0000000..5e7d0cc --- /dev/null +++ b/source/utils/ansi.py @@ -0,0 +1,113 @@ +# Copyright Jonathan Hartley 2013. BSD 3-Clause license, see LICENSE file. +''' +This module generates ANSI character codes to printing colors to terminals. +See: http://en.wikipedia.org/wiki/ANSI_escape_code +''' +import sys +import os + +CSI = '\033[' +OSC = '\033]' +BEL = '\a' +is_a_tty = sys.stderr.isatty() and os.name == 'posix' + + +def code_to_chars(code): + return CSI + str(code) + 'm' if is_a_tty else '' + + +def set_title(title): + return OSC + '2;' + title + BEL + + +def clear_screen(mode=2): + return CSI + str(mode) + 'J' + + +def clear_line(mode=2): + return CSI + str(mode) + 'K' + + +class AnsiCodes(object): + def __init__(self): + # the subclasses declare class attributes which are numbers. + # Upon instantiation we define instance attributes, which are the same + # as the class attributes but wrapped with the ANSI escape sequence + for name in dir(self): + if not name.startswith('_'): + value = getattr(self, name) + setattr(self, name, code_to_chars(value)) + + +class AnsiCursor(object): + def UP(self, n=1): + return CSI + str(n) + 'A' + + def DOWN(self, n=1): + return CSI + str(n) + 'B' + + def FORWARD(self, n=1): + return CSI + str(n) + 'C' + + def BACK(self, n=1): + return CSI + str(n) + 'D' + + def POS(self, x=1, y=1): + return CSI + str(y) + ';' + str(x) + 'H' + + +class AnsiFore(AnsiCodes): + BLACK = 30 + RED = 31 + GREEN = 32 + YELLOW = 33 + BLUE = 34 + MAGENTA = 35 + CYAN = 36 + WHITE = 37 + RESET = 39 + + # These are fairly well supported, but not part of the standard. + LIGHTBLACK_EX = 90 + LIGHTRED_EX = 91 + LIGHTGREEN_EX = 92 + LIGHTYELLOW_EX = 93 + LIGHTBLUE_EX = 94 + LIGHTMAGENTA_EX = 95 + LIGHTCYAN_EX = 96 + LIGHTWHITE_EX = 97 + + +class AnsiBack(AnsiCodes): + BLACK = 40 + RED = 41 + GREEN = 42 + YELLOW = 43 + BLUE = 44 + MAGENTA = 45 + CYAN = 46 + WHITE = 47 + RESET = 49 + + # These are fairly well supported, but not part of the standard. + LIGHTBLACK_EX = 100 + LIGHTRED_EX = 101 + LIGHTGREEN_EX = 102 + LIGHTYELLOW_EX = 103 + LIGHTBLUE_EX = 104 + LIGHTMAGENTA_EX = 105 + LIGHTCYAN_EX = 106 + LIGHTWHITE_EX = 107 + + +class AnsiStyle(AnsiCodes): + BRIGHT = 1 + DIM = 2 + NORMAL = 22 + RESET_ALL = 0 + + +Fore = AnsiFore() +Back = AnsiBack() +Style = AnsiStyle() +Cursor = AnsiCursor() diff --git a/source/utils/comfy_node_action.py b/source/utils/comfy_node_action.py new file mode 100644 index 0000000..82d090c --- /dev/null +++ b/source/utils/comfy_node_action.py @@ -0,0 +1,44 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# ComfyUI Node actions +import logging +# ComfyUI imports +try: + from server import PromptServer + with_comfy = True +except Exception: + with_comfy = False +# Local imports +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.comfy_node_action") + + +def send_node_action(action: str, arg1: str = None, arg2: str = None, sid: str = None): + """ + Sends a node action event to the ComfyUI client. + + Args: + action (str): Action to be performed. + arg1 (str): First argument + arg2 (str): Second argument + sid (str, optional): The session ID of the client to send to. + If None, broadcasts to all clients. Defaults to None. + """ + if not with_comfy: + return + try: + PromptServer.instance.send_sync( + "set-audioseparation-node", # This is our custom event name + { + 'action': action, + 'arg1': arg1, + 'arg2': arg2 + }, + sid + ) + except Exception as e: + logger.error(f"when trying to use ComfyUI PromptServer: {e}") diff --git a/source/utils/comfy_notification.py b/source/utils/comfy_notification.py new file mode 100644 index 0000000..657c7a6 --- /dev/null +++ b/source/utils/comfy_notification.py @@ -0,0 +1,46 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# ComfyUI Toast API messages +# Original code from Gemini 2.5 Pro, which was really outdated +# Took ideas from Easy Use nodes and looking at ComfyUI code +import logging +# ComfyUI imports +try: + from server import PromptServer + with_comfy = True +except Exception: + with_comfy = False +# Local imports +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.comfy_notification") + + +def send_toast_notification(message: str, summary: str = "Warning", severity: str = "warn", sid: str = None): + """ + Sends a toast notification event to the ComfyUI client. + + Args: + message (str): The message content of the toast. + severity (str): The type of toast. Can be 'success' | 'info' | 'warn' | 'error' | 'secondary' | 'contrast' + summary (str): Short explanation + sid (str, optional): The session ID of the client to send to. + If None, broadcasts to all clients. Defaults to None. + """ + if not with_comfy: + return + try: + PromptServer.instance.send_sync( + "set-audioseparation-toast", # This is our custom event name + { + 'message': message, + 'summary': summary, + 'severity': severity + }, + sid + ) + except Exception as e: + logger.error(f"when trying to use ComfyUI PromptServer: {e}") diff --git a/source/utils/downloader.py b/source/utils/downloader.py new file mode 100644 index 0000000..f057be1 --- /dev/null +++ b/source/utils/downloader.py @@ -0,0 +1,206 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Model downloader w/TQDM and ComfyUI progress +# Original code from Gemini 2.5 Pro +import logging +import os +# Requests is better than the core Python urllib, and is a really common package +# But we don't really need it. Lets make it optional: +try: + import requests + with_requests = True +except Exception: + with_requests = False +import urllib +from tqdm import tqdm +# ComfyUI imports +try: + import comfy.utils + with_comfy = True +except Exception: + with_comfy = False +# Local imports +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.downloader") + + +def download_model_requests(url: str, save_dir: str, file_name: str): + """ + Downloads a file from a URL with progress bars for both console and ComfyUI. + + Args: + url (str): The direct download URL for the file. + save_dir (str): The directory where the file will be saved. + file_name (str): The name of the file to be saved on disk. + """ + full_path = os.path.join(save_dir, file_name) + + # Ensure the save directory exists + os.makedirs(save_dir, exist_ok=True) + try: + # Use a streaming request to handle large files and get content length + with requests.get(url, stream=True, timeout=10) as r: + r.raise_for_status() # Raise an exception for bad status codes (4xx or 5xx) + + # Get total file size from headers + total_size_in_bytes = int(r.headers.get('content-length', 0)) + block_size = 1024 # 1 Kibibyte + + # --- Setup Progress Bars --- + # Console progress bar using tqdm + progress_bar_console = tqdm( + total=total_size_in_bytes, + unit='iB', + unit_scale=True, + desc=f"Downloading {file_name}" + ) + + # ComfyUI progress bar + progress_bar_ui = comfy.utils.ProgressBar(total_size_in_bytes) if with_comfy else None + + # --- Download Loop --- + downloaded_size = 0 + with open(full_path, 'wb') as f: + for chunk in r.iter_content(chunk_size=block_size): + if chunk: # filter out keep-alive new chunks + chunk_size = len(chunk) + + # Update console progress bar + progress_bar_console.update(chunk_size) + + # Update ComfyUI progress bar + downloaded_size += chunk_size + if progress_bar_ui: + progress_bar_ui.update(chunk_size) # ProgressBar takes absolute value, but update is incremental + + # Write chunk to file + f.write(chunk) + + # --- Cleanup --- + progress_bar_console.close() + + # Final check to see if download was complete + if total_size_in_bytes != 0 and progress_bar_console.n != total_size_in_bytes: + logger.error("Download failed: Size mismatch.") + # Optional: remove partial file + # os.remove(full_path) + raise IOError(f"Download failed for {file_name}. Expected {total_size_in_bytes} but got " + f"{progress_bar_console.n}") + + return full_path + + except requests.exceptions.RequestException as e: + logger.error(f"Network error while downloading {file_name}: {e}") + # Clean up partial file if it exists + if os.path.exists(full_path): + try: + os.remove(full_path) + except OSError: + pass + raise + except Exception as e: + logger.error(f"An error occurred during download: {e}") + if os.path.exists(full_path): + try: + os.remove(full_path) + except OSError: + pass + raise + + +# A simple version implemented using the Python urllib +class Downloader: + def __init__(self, model_path, model_name): + self.model_path = model_path + self.model_name = model_name + self.model_full_name = os.path.join(self.model_path, self.model_name) + # Ensure the directory for the model_path exists before __init__ if used elsewhere + # or create it at the start of download_model + + # A TQDM helper class for urlretrieve reporthook + # This is a common pattern for this use case. + class TqdmUpTo(tqdm): + """ + Provides `update_to(block_num, block_size, total_size)` + and updates the TQDM bar. + """ + def __init__(self, unit, unit_scale, unit_divisor, miniters, desc): + super().__init__(unit=unit, unit_scale=unit_scale, unit_divisor=unit_divisor, miniters=miniters, desc=desc) + self.ui_bar = None + self.total = None + + def update_to(self, block_num=1, block_size=1, total_size=None): + """ + block_num : int, optional + Number of blocks transferred so far [default: 1]. + block_size : int, optional + Size of each block (in tqdm units) [default: 1]. + total_size : int, optional + Total size (in tqdm units). If [default: None] remains unchanged. + """ + if total_size is not None and self.total is None: + self.total = total_size + # ComfyUI progress bar + if self.ui_bar is None and with_comfy: + self.ui_bar = comfy.utils.ProgressBar(total_size) + # self.update() will take the *difference* from the last call. + # So we pass the number of new blocks * block_size. + # Since block_num is cumulative, we calculate the new amount. + chunk_size = block_num * block_size - self.n + self.update(chunk_size) # self.n is current progress + if self.ui_bar: + self.ui_bar.update(chunk_size) # ProgressBar takes absolute value, but update is incremental + + def download_model(self, url: str): + try: + # Ensure the directory exists + # Use or '.' for current dir if dirname is empty + os.makedirs(self.model_path or '.', exist_ok=True) + + # Get filename for tqdm description + filename = self.model_name + + # Use TqdmUpTo as a context manager + with self.TqdmUpTo(unit='iB', unit_scale=True, unit_divisor=1024, miniters=1, + desc=f"Downloading {filename}") as t: + # urlretrieve(url, filename=None, reporthook=None, data=None) + # reporthook is called with (block_num, block_size, total_size) + urllib.request.urlretrieve(url, self.model_full_name, reporthook=t.update_to) + # The 'with' statement ensures t.close() is called. + + return filename + + except urllib.error.URLError as e: # More specific exception for network issues + # Clean up partially downloaded file if an error occurs + if os.path.exists(self.model_full_name): + os.remove(self.model_full_name) + raise Exception(f"An error occurred while downloading the model (URL Error): {e.reason} from {url}") + + except Exception as e: + # Clean up partially downloaded file if an error occurs + if os.path.exists(self.model_full_name): + os.remove(self.model_full_name) + raise Exception(f"An unexpected error occurred while downloading the model: {e}") + + +def download_model_urllib(url: str, save_dir: str, file_name: str): + return Downloader(save_dir, file_name).download_model(url) + + +def download_model(url: str, save_dir: str, file_name: str, force_urllib: bool = False): + logger.info(f"Downloading model: {file_name}") + logger.info(f"Source URL: {url}") + full_name = os.path.join(save_dir, file_name) + logger.info(f"Destination: {full_name}") + + if with_requests and not force_urllib: + download_model_requests(url, save_dir, file_name) + else: + download_model_urllib(url, save_dir, file_name) + + logger.info(f"Successfully downloaded {full_name}") + return full_name diff --git a/source/utils/load_audio.py b/source/utils/load_audio.py new file mode 100644 index 0000000..3b8e080 --- /dev/null +++ b/source/utils/load_audio.py @@ -0,0 +1,49 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Audio load helper +# Original code from Gemini 2.5 Pro +import logging +import torch +import torchaudio +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.load_audio") + + +def audio_get_channels(waveform): + dim_c = 0 if waveform.ndim == 2 else 1 + return waveform.shape[dim_c] + + +def force_stereo(waveform): + dim_c = 0 if waveform.ndim == 2 else 1 + if waveform.shape[dim_c] == 1: + logger.debug("Audio is mono, converting to fake stereo.") + return torch.cat([waveform, waveform], dim=dim_c) + return waveform + + +def force_sample_rate(waveform, orig_freq, new_freq): + logger.debug(f"Resampling from {orig_freq} Hz to {new_freq} Hz.") + resampler = torchaudio.transforms.Resample(orig_freq=orig_freq, new_freq=new_freq) + return resampler(waveform) + + +def load_audio(file_path, force_sr=None, force_stereo=False): + """ Loads an audio file, optionally converts it to stereo float, and resamples to force_sr. """ + logger.info(f"🎵 Loading audio file: {file_path}") + try: + waveform, sample_rate = torchaudio.load(file_path, normalize=True) + # Ensure stereo + if force_stereo and audio_get_channels(waveform) == 1: + waveform = force_stereo(waveform) + # Ensure 44.1 kHz or other S/R + if force_sr is not None and sample_rate != force_sr: + waveform = force_sample_rate(waveform, sample_rate, force_sr) + return waveform, sample_rate + except Exception as e: + logger.error(f"💥 Failed to load audio file: {e}") + raise diff --git a/source/utils/load_class.py b/source/utils/load_class.py new file mode 100644 index 0000000..3b8faa5 --- /dev/null +++ b/source/utils/load_class.py @@ -0,0 +1,62 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# This helper is used to load a class from an arbitrary file +# Gemini 2.5 Pro code +import importlib +import logging +import os +import sys +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.load_class") + + +# Helper to dynamically import the target PyTorch model class +def import_model_class(location_string: str): + """ + Dynamically imports a PyTorch model class from a file path and class name. + + The location_string is expected to be in the format: + 'path/to/your/file.py:ClassName' + """ + module_dir = None + try: + # 1. Split the input string into a file path and a class name + filepath, class_name = location_string.split(':') + + # Check if the file exists before proceeding + if not os.path.exists(filepath): + logger.error(f"File not found at '{filepath}'.") + sys.exit(1) + + # 2. Get the directory and the module name from the file path + module_dir, module_file = os.path.split(filepath) + module_name = os.path.splitext(module_file)[0] + + # Add the directory to sys.path to allow Python to find it + # Add it to the beginning to ensure it's checked first + sys.path.insert(0, module_dir) + + # 3. Import the module + logger.info(f"Importing module '{module_name}' from '{module_dir}'...") + module = importlib.import_module(module_name) + + # 4. Get the class from the imported module + model_class = getattr(module, class_name) + + except (ValueError, ImportError, AttributeError, FileNotFoundError) as e: + logger.error(f"Could not import model class from '{location_string}'.") + logger.error("Please ensure the format is 'path/to/file.py:ClassName'.") + logger.error(f"Original error: {e}") + sys.exit(1) + + finally: + # 5. Clean up by removing the path we added. + # This is crucial to avoid polluting the user's environment. + if module_dir is not None and module_dir in sys.path: + sys.path.pop(0) + + logger.info(f"Successfully imported class '{class_name}'.") + return model_class diff --git a/source/utils/load_onnx.py b/source/utils/load_onnx.py new file mode 100644 index 0000000..a2c59d4 --- /dev/null +++ b/source/utils/load_onnx.py @@ -0,0 +1,59 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# ONNX model loader helper +# Original code from Gemini 2.5 Pro +import logging +try: + import onnxruntime as ort + with_onnx = True +except Exception: + with_onnx = False +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.load_onnx") + +if with_onnx: + import torch + + class ONNXWrapper: + """ + A wrapper class for an ONNX Runtime InferenceSession to provide a + PyTorch-like __call__ interface. + """ + def __init__(self, session: ort.InferenceSession, device: torch.device): + self.session = session + self.device = device + # Get the name of the input tensor from the model's graph + self.input_name = self.session.get_inputs()[0].name + + def __call__(self, input_tensor: torch.Tensor): + """ + Performs inference using the ONNX session. + + Args: + input_tensor: A PyTorch tensor already on the correct device. + + Returns: + A PyTorch tensor on the same device as the input. + """ + # 1. Convert the input PyTorch tensor to a CPU NumPy array + input_numpy = input_tensor.cpu().numpy() + # 2. Run the ONNX session + result_numpy = self.session.run(None, {self.input_name: input_numpy})[0] + # 3. Convert the output NumPy array back to a PyTorch tensor on the original device + result_tensor = torch.from_numpy(result_numpy).to(self.device) + + return result_tensor + + def load_onnx(model_path, device): + logger.info("Loading ONNX model for runtime inference...") + providers = ['CUDAExecutionProvider' if 'cuda' in str(device) else 'CPUExecutionProvider'] + session = ort.InferenceSession(model_path, providers=providers) + model_w = ONNXWrapper(session, device) + return model_w +else: + def load_onnx(model_path, device): + raise ValueError("No ONNX support, please install `onnxruntime`") diff --git a/source/utils/load_safetensors.py b/source/utils/load_safetensors.py new file mode 100644 index 0000000..4d19dc6 --- /dev/null +++ b/source/utils/load_safetensors.py @@ -0,0 +1,32 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Model loader helper +# Original code from Gemini 2.5 Pro +import logging +from safetensors.torch import load_file +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.load_safetensors") + + +def load_safetensors(model_path, model_run, device): + logger.info("Loading PyTorch model from .safetensors file...") + # 1. Load the state_dict from the file, EXPLICITLY forcing all tensors onto the CPU. + state_dict = load_file(model_path, device="cpu") + # 2. Load the CPU state_dict into the CPU model. This is now a safe operation. + try: + missing_keys, unexpected_keys = model_run.load_state_dict(state_dict, strict=False) + if missing_keys: + logger.warning(f"Missing keys in state_dict for model_run: {missing_keys}") + if unexpected_keys: + logger.warning(f"Unexpected keys in state_dict for model_run: {unexpected_keys}") + if not missing_keys and not unexpected_keys: + logger.debug("All keys matched successfully.") + except RuntimeError as e: + logger.error(f"RuntimeError during model_run.load_state_dict: {e}") + logger.error("This might indicate a mismatch between saved weights and model architecture.") + raise + return model_run diff --git a/source/utils/logger.py b/source/utils/logger.py new file mode 100644 index 0000000..a31d7d2 --- /dev/null +++ b/source/utils/logger.py @@ -0,0 +1,105 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +import os +import sys +import logging +from .misc import NODES_NAME, NODES_DEBUG_VAR + +no_colorama = False +try: + from colorama import init as colorama_init, Fore, Back, Style +except ImportError: + no_colorama = True +# If colorama isn't installed use an ANSI basic replacement +if no_colorama: + from .ansi import Fore, Back, Style # noqa: F811 +else: + colorama_init() + +# Used for tools +standalone_mode = False + +white = Fore.WHITE + Style.BRIGHT +yellow = Fore.YELLOW + Style.BRIGHT +red = Fore.RED + Style.BRIGHT +red_alarm = Fore.RED + Back.WHITE + Style.BRIGHT +cyan = Fore.CYAN + Style.BRIGHT +reset = Style.RESET_ALL +# format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s " +# "(%(filename)s:%(lineno)d)" +format = f"[{NODES_NAME} %(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)" +format_simple = f"[{NODES_NAME}] %(message)s" +FORMATS = { + logging.DEBUG: cyan + format + reset, + logging.INFO: white + format_simple + reset, + logging.WARNING: yellow + format + reset, + logging.ERROR: red + format + reset, + logging.CRITICAL: red_alarm + format + reset +} +format = "[%(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)" +format_simple = "%(message)s" +if not sys.stdout.isatty(): + white = yellow = red = red_alarm = cyan = reset = "" +FORMATS_STANDALONE = { + logging.DEBUG: cyan + format + reset, + logging.INFO: white + format_simple + reset, + logging.WARNING: yellow + format + reset, + logging.ERROR: red + format + reset, + logging.CRITICAL: red_alarm + format + reset +} + + +class CustomFormatter(logging.Formatter): + """Logging Formatter to add colors""" + + def __init__(self): + super(logging.Formatter, self).__init__() + + def format(self, record): + formats = FORMATS_STANDALONE if standalone_mode else FORMATS + log_fmt = formats.get(record.levelno) + formatter = logging.Formatter(log_fmt) + return formatter.format(record) + + +# Create a new logger +logger = logging.getLogger(NODES_NAME) +logger.propagate = False + +# Add handler if we don't have one. +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(CustomFormatter()) + logger.addHandler(handler) + +# ###################### +# Logger setup +# ###################### +# 1. Determine the ComfyUI global log level (influenced by --verbose) +main_logger = logger +comfy_root_logger = logging.getLogger('comfy') +effective_comfy_level = logging.getLogger().getEffectiveLevel() +# 2. Check our custom environment variable for more verbosity +try: + nodes_debug_env = int(os.environ.get(NODES_DEBUG_VAR, "0")) +except ValueError: + nodes_debug_env = 0 +# 3. Set node's logger level +if nodes_debug_env: + main_logger.setLevel(logging.DEBUG - (nodes_debug_env - 1)) + final_level_str = f"DEBUG (due to {NODES_DEBUG_VAR}={nodes_debug_env})" +else: + main_logger.setLevel(effective_comfy_level) + final_level_str = logging.getLevelName(effective_comfy_level) + " (matching ComfyUI global)" +_initial_setup_logger = logging.getLogger(NODES_NAME + ".setup") # A temporary logger for this message +_initial_setup_logger.debug(f"{NODES_NAME} logger level set to: {final_level_str}") + + +def logger_set_standalone(args): + verbose = args.verbose + global main_logger + main_logger.setLevel(logging.DEBUG - (verbose - 1) if verbose else logging.INFO) + global standalone_mode + standalone_mode = True diff --git a/source/utils/misc.py b/source/utils/misc.py new file mode 100644 index 0000000..1177988 --- /dev/null +++ b/source/utils/misc.py @@ -0,0 +1,18 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +import logging + +NODES_NAME = "AudioSeparation" +NODES_DEBUG_VAR = NODES_NAME.upper() + "_NODES_DEBUG" + + +def debugl(logger, level, msg): + if logger.getEffectiveLevel() <= logging.DEBUG - (level - 1): + logger.debug(msg) + + +def cli_add_verbose(parser): + parser.add_argument('-v', '--verbose', action='count', default=0, + help="Enable verbose output to see details of the process.") diff --git a/source/utils/save_audio.py b/source/utils/save_audio.py new file mode 100644 index 0000000..466fcca --- /dev/null +++ b/source/utils/save_audio.py @@ -0,0 +1,41 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Audio save helper +# Original code from Gemini 2.5 Pro +import logging +import os +import torchaudio +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.save_audio") + + +def save_audio(tensor, sample_rate, file_path, output_format): + """ + Saves a tensor as an audio file, using the most basic and compatible + torchaudio.save signature to avoid all version-specific errors. + """ + logger.info(f"💾 Saving audio to: {file_path}") + + output_dir = os.path.dirname(file_path) + if output_dir and not os.path.exists(output_dir): + os.makedirs(output_dir, exist_ok=True) + + try: + # The most compatible signature is simply: + # torchaudio.save(filepath, src, sample_rate, format) + # We pass the format string directly. The ffmpeg backend will use + # a reasonable default quality for MP3 encoding. + torchaudio.save(file_path, tensor.cpu(), sample_rate, format=output_format.lower()) + + logger.info("✅ Save complete.") + except Exception as e: + if "ffmpeg" in str(e).lower() and "Unknown encoder" not in str(e): + logger.error("💥 Failed to save audio file. This might be because the 'ffmpeg' backend is not available.") + logger.error("Please ensure FFmpeg is installed and accessible in your system's PATH.") + else: + logger.error(f"💥 Failed to save audio file: {e}") + raise diff --git a/source/utils/torch.py b/source/utils/torch.py new file mode 100644 index 0000000..70a9080 --- /dev/null +++ b/source/utils/torch.py @@ -0,0 +1,111 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: CC BY-NC-SA 4.0 +# Project: ComfyUI-Float_Optimized +import contextlib # For context manager +import logging +import torch +try: + import comfy.model_management as mm + with_comfy = True +except Exception: + with_comfy = False +from .misc import NODES_NAME + +logger = logging.getLogger(f"{NODES_NAME}.torch") + + +def get_torch_device_options(): + # We always have CPU + default = "cpu" + options = [default] + # Do we have CUDA? + if torch.cuda.is_available(): + default = "cuda" + options.append(default) + for i in range(torch.cuda.device_count()): + options.append(f"cuda:{i}") # Specific CUDA devices + # Is this a Mac? + if torch.backends.mps.is_available() and torch.backends.mps.is_built(): + options.append("mps") + if default == "cpu": + default = "mps" + return options, default + + +# ################################################################################## +# # Helper for inference (Target device, offload, eval, no_grad and cuDNN Benchmark) +# ################################################################################## + +@contextlib.contextmanager +def model_to_target(model): + """ + Consolidated context manager for model device placement and inference state. + + - Moves the model to its designated `model.target_device`. + - Sets `torch.backends.cudnn.benchmark` based on `model.cudnn_benchmark_setting` if available. + - Sets the model to `eval()` mode. + - Wraps the operation in a `torch.no_grad()` context. + - Offloads the model to the CPU (`mm.unet_offload_device()`) afterwards. + """ + if not isinstance(model, torch.nn.Module): + with torch.no_grad(): + yield # The code inside the 'with' statement runs here + return + + # 1. Determine target device from the model object + try: + target_device = model.target_device + assert isinstance(target_device, torch.device) + except Exception as e: + logger.warning(f"model_to_target: Could not get 'target_device' from model ({e}). " + "Defaulting to model's current device.") + target_device = next(model.parameters()).device + + # 2. Get CUDNN benchmark setting from the model object (optional) + # Use hasattr as this is an optional setting that not all models might have. + cudnn_benchmark_enabled = None # Default is to keep the current setting + if hasattr(model, 'cudnn_benchmark_setting'): + cudnn_benchmark_enabled = model.cudnn_benchmark_setting + + original_device = next(model.parameters()).device + original_cudnn_benchmark_state = None + is_cuda_target = target_device.type == 'cuda' + + try: + # 3. Manage cuDNN benchmark state + if (cudnn_benchmark_enabled is not None and is_cuda_target and hasattr(torch.backends, 'cudnn') and + torch.backends.cudnn.is_available()): + if torch.backends.cudnn.benchmark != cudnn_benchmark_enabled: + original_cudnn_benchmark_state = torch.backends.cudnn.benchmark + torch.backends.cudnn.benchmark = cudnn_benchmark_enabled + logger.debug(f"Temporarily set cuDNN benchmark to {torch.backends.cudnn.benchmark}") + + # 4. Move model to target device if not already there + if original_device != target_device: + logger.debug(f"Moving model from `{original_device}` to target device `{target_device}` for inference.") + model.to(target_device) + + # 5. Set to eval mode and disable gradients for the operation + model.eval() + with torch.no_grad(): + yield # The code inside the 'with' statement runs here + + finally: + # 6. Restore original cuDNN benchmark state + if original_cudnn_benchmark_state is not None: + # This check is sufficient because it will only be not None if we set it inside the try block + torch.backends.cudnn.benchmark = original_cudnn_benchmark_state + logger.debug(f"Restored cuDNN benchmark to {original_cudnn_benchmark_state}") + + # 7. Offload model back to CPU + if with_comfy: + offload_device = mm.unet_offload_device() + current_device_after_yield = next(model.parameters()).device + if current_device_after_yield != offload_device: + logger.debug(f"Offloading model from `{current_device_after_yield}` to offload device `{offload_device}`.") + model.to(offload_device) + # Clear cache if we were on a CUDA device + if 'cuda' in str(current_device_after_yield): + torch.cuda.empty_cache() diff --git a/tool/bootstrap/__init__.py b/tool/bootstrap/__init__.py new file mode 100644 index 0000000..fae4571 --- /dev/null +++ b/tool/bootstrap/__init__.py @@ -0,0 +1,55 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Why such a complex thing? +# Python imports are broken by design, and here we hit a huge limitation: +# 1. ComfyUI nodes MUST use relative imports, if you don't do it things like "import utils" becomes ambiguous +# You might think this can be overcome polluting the sys.path, but this isn't true. If ComfyUI, or some +# other node, already imported an module named "utils" you'll get the already imported module, not the one +# you want. And polluting sys.path can make other nodes, or ComfyUI itself, import the wrong module when +# they use a "non top-level import". +# Conclusion: The only safe way to import from a ComfyUI node is by using relative imports. +# 2. We have tools in the "tool" subdir, they are intended to run as standalone scripts. They can't use +# relative imports to access the node sub modules because you'll hit the error "ImportError: attempted +# relative import with no known parent package". So imports in tools MUST be absolute. This can be solved +# adding the node root to sys.path. Here we are not polluting the sys.path because we are the top-level. +# Conclusion: This script is used to solve adding the correct path to sys.path +# 3. As you MUST use relative imports in the node sub-modules, when a sub-module depends on another sub-module +# it will do something like "from ..XXXX", this is OK when all is relative. But tools are using absolute +# imports, so you'll hit the error "ImportError: attempted relative import beyond top-level package" +# Conclusion: All sub-modules must be wrapped by an umbrella sub-module, this is what "source" is. +# +# Structure: +# Node-name/ <-- is a module from ComfyUI point of view, dynamically imported using importlib +# \-- __init__.py <-- needed because we are a module, relative imports +# \-- nodes.py <-- relative imports +# | +# \-- source/ <-- umbrella package, just to solve relative imports +# | \-- __init__.py +# | | +# | \-- utils/ <-- a submodule, uses relative imports +# | | \-- __init__.py +# | | \-- misc.py +# | | +# | \-- db/ <-- another submodule, uses relative imports +# | \-- __init__.py +# | \-- models_db.py <-- Must use relative imports, can access ..utils.misc +# | +# \-- tool/ +# \-- batch_convert.py <-- a tool, must use absolute import of "source" +# | +# \-- bootstrap/ +# \-- __init__.py <-- THIS file, adds to sys.path to make "source" available + +# --- BOOTSTRAP: Make the script aware of the project root --- +import os +import sys +# 1. Get the absolute path of the current script's directory (tool/) +script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +# 2. Get the project root by going one directory up (MyAwesomeNode/) +project_root = os.path.dirname(script_dir) +# 3. Add the project root to the Python path +if project_root not in sys.path: + sys.path.insert(0, project_root) diff --git a/tool/demix.py b/tool/demix.py new file mode 100644 index 0000000..baef733 --- /dev/null +++ b/tool/demix.py @@ -0,0 +1,128 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Tool to demix audio +# Run it using: python tool/demix.py -m HASH AUDIO +import argparse +import os +import sys +import torch +# Local imports +import bootstrap # noqa: F401 +from source.db.models_db import ModelsDB, cli_add_models_and_db +from source.inference.demixer import get_demixer +from source.utils.logger import main_logger, logger_set_standalone +from source.utils.load_audio import load_audio +from source.utils.save_audio import save_audio +from source.utils.misc import cli_add_verbose + +BANNER = "🎵 MDX-Net Audio Separation Tool 🎵" + + +# --- Main Demixing Logic --- +def demix(d, args): + main_logger.info("🚀 Starting audio separation process...") + device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + main_logger.info(f"💻 Using device: {device}") + + # --- Load and Prepare Audio --- + waveform, sr = load_audio(args.input_file, force_sr=44100, force_stereo=True) + + # --- Load Model --- + demixer = get_demixer(d, device, args.models_dir) + + # --- Do inference in chunks --- + wavs = demixer(waveform, args.segments) + + # --- Save outputs --- + base, ext = os.path.splitext(args.input_file) + out_ext = ext if not args.format else '.' + args.format.lower() + out_format = args.format or out_ext[1:] + if args.save_main: + out_path = f"{args.out_base or base}_{wavs[0]['stem']}{out_ext}" + save_audio(wavs[0]['waveform'], wavs[0]['sample_rate'], out_path, out_format) + + if args.save_complement: + for wav in wavs[1:]: + out_path = f"{args.out_base or base}_{wav['stem']}{out_ext}" + save_audio(wav['waveform'], wav['sample_rate'], out_path, out_format) + + main_logger.info("🎉 All operations complete!") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser( + description=BANNER, + formatter_class=argparse.ArgumentDefaultsHelpFormatter + ) + + # --- Input options --- + parser.add_argument('input_file', type=str, nargs='?', default=None, + help="Path to the input audio file (wav, mp3, flac). Required unless --list is used.") + parser.add_argument('-m', '--model', type=str, default="499a6a6bf9da6d330235a1576007ddc0", + help="Hash or name for known model (use --list to know all). " + "Or path to the model file (.safetensors or .onnx), which must be known.") + cli_add_models_and_db(parser) + + # --- Output options --- + parser.add_argument('--no_main', dest='save_main', action='store_false', + help="Do not save the main separated stem.") + parser.add_argument('--save_complement', action='store_true', + help="Save the complement stem (input - main).") + parser.add_argument('--out_base', type=str, default=None, + help="Base for the output path. No extension here, we will add the name of the stem and extension") + parser.add_argument('--format', type=str, default=None, choices=['wav', 'flac', 'mp3'], + help="Output audio format. Defaults to input format.") + + # --- Control Arguments --- + parser.add_argument('--segments', type=int, default=1, + help="How many audio segments to process at once") + parser.add_argument('-l', '--list', action='store_true', help="Show available models") + cli_add_verbose(parser) + + parser.set_defaults(save_main=True) + args = parser.parse_args() + logger_set_standalone(args) + main_logger.info(BANNER) + + # Sanity check + if not args.list: + if not args.input_file: + main_logger.error("💥 The following arguments are required: input_file") + sys.exit(2) + if not args.save_main and not args.save_complement: + main_logger.error("💥 Nothing to save! Please don't specify --no_main (default) or use --save_complement.") + sys.exit(2) + + # Check what we have + db = ModelsDB(args.models_dir, args.json_file) + models = db.get_filtered() + + # Just list models + if args.list: + main_logger.info("\n--- Available models ---\n") + known = models.get_display_names(clean=True) + for m in known: + main_logger.info(f"`{m}` {models.get_by_display_name(m)['hash']}") + main_logger.info(f"\n{len(known)} models known.") + sys.exit(0) + + # Check we have a sound file + if not os.path.isfile(args.input_file): + main_logger.error(f"💥 `{args.input_file}` is not a file.") + sys.exit(4) + + # Look for the selected model + d = models.get(args.model) + if d is None: + main_logger.error(f"💥 Unknown model `{args.model}`.") + sys.exit(3) + try: + main_logger.info(f"📂 Using model from `{d['model_path']}`") + except KeyError: + # Needs download + pass + + demix(d, args)