[Added] Working node and CLI tool
This commit is contained in:
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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'])
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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}")
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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`")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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.")
|
||||
@@ -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
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user