[Added] Working node and CLI tool

This commit is contained in:
Salvador E. Tropea
2025-07-02 08:33:52 -03:00
parent b0a1e03326
commit e0b2bde974
29 changed files with 2375 additions and 0 deletions
View File
+113
View File
@@ -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()
+44
View File
@@ -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}")
+46
View File
@@ -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}")
+206
View File
@@ -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
+49
View File
@@ -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
+62
View File
@@ -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
+59
View File
@@ -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`")
+32
View File
@@ -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
+105
View File
@@ -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
+18
View File
@@ -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.")
+41
View File
@@ -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
+111
View File
@@ -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()