Refactored node registration to handle any custom node type by extracting class definitions from source files and injecting wrapper code with proper indentation. Replaces hardcoded LTX module registration while maintaining identical functionality. Successfully tested with LTXVLoader node.
316 lines
13 KiB
Python
316 lines
13 KiB
Python
import time
|
|
import copy
|
|
import torch
|
|
import comfy.model_management
|
|
import os
|
|
from pathlib import Path # Add this import
|
|
import importlib.util
|
|
import logging
|
|
import folder_paths
|
|
|
|
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
|
logging.info("MultiGPU: Initialization started")
|
|
|
|
current_device = comfy.model_management.get_torch_device()
|
|
logging.info(f"MultiGPU: Initial device {current_device}")
|
|
|
|
def get_torch_device_patched():
|
|
if (
|
|
not torch.cuda.is_available()
|
|
or comfy.model_management.cpu_state == comfy.model_management.CPUState.CPU
|
|
or "cpu" in str(current_device).lower()
|
|
):
|
|
return torch.device("cpu")
|
|
return torch.device(current_device)
|
|
|
|
comfy.model_management.get_torch_device = get_torch_device_patched
|
|
|
|
def override_class(cls):
|
|
class NodeOverride(cls):
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
|
devices = ["cpu"] + [f"cuda:{i}" for i in range(torch.cuda.device_count())]
|
|
inputs["required"]["device"] = (devices,)
|
|
return inputs
|
|
|
|
CATEGORY = "multigpu"
|
|
FUNCTION = "override"
|
|
|
|
def override(self, *args, device, **kwargs):
|
|
global current_device
|
|
current_device = device
|
|
fn = getattr(super(), cls.FUNCTION)
|
|
return fn(*args, **kwargs)
|
|
|
|
return NodeOverride
|
|
|
|
NODE_CLASS_MAPPINGS = {}
|
|
|
|
def register_module_new(module_path, target_nodes):
|
|
module_dir = os.path.join("custom_nodes", module_path)
|
|
if not os.path.isdir(module_dir):
|
|
logging.warning(f"MultiGPU: Module directory {module_path} not found.")
|
|
return
|
|
|
|
node_classes = {}
|
|
|
|
for node_name in target_nodes:
|
|
node_definition_block = None
|
|
for filename in os.listdir(module_dir):
|
|
if filename.endswith(".py"):
|
|
filepath = os.path.join(module_dir, filename)
|
|
try:
|
|
with open(filepath, 'r') as f:
|
|
lines = f.readlines()
|
|
|
|
start_index = -1
|
|
second_def_index = -1
|
|
def_count = 0
|
|
in_class = False
|
|
|
|
for i, line in enumerate(lines):
|
|
if line.strip().startswith(f"class {node_name}"):
|
|
start_index = i
|
|
in_class = True
|
|
elif in_class and line.strip().startswith("def "):
|
|
def_count += 1
|
|
if def_count == 2:
|
|
second_def_index = i + 1
|
|
break
|
|
|
|
if start_index != -1 and second_def_index != -1:
|
|
definition_lines = lines[start_index:second_def_index]
|
|
last_func_indent = len(definition_lines[-1]) - len(definition_lines[-1].lstrip())
|
|
wrapper_code = [
|
|
" " * (last_func_indent + 4) + "from nodes import NODE_CLASS_MAPPINGS\n",
|
|
" " * (last_func_indent + 4) + f"original_loader = NODE_CLASS_MAPPINGS[\"{node_name}\"]()\n",
|
|
" " * (last_func_indent + 4) + "return original_loader.load(ckpt_name, dtype)\n"
|
|
]
|
|
node_definition_block = "".join(definition_lines) + "".join(wrapper_code)
|
|
break
|
|
|
|
except Exception as e:
|
|
logging.error(f"MultiGPU: Error processing file {filepath}: {str(e)}")
|
|
continue
|
|
|
|
if node_definition_block:
|
|
namespace = {}
|
|
exec("import folder_paths", namespace)
|
|
exec("from pathlib import Path", namespace)
|
|
|
|
try:
|
|
exec(node_definition_block, namespace)
|
|
node_classes[node_name] = namespace[node_name]
|
|
logging.info(f"MultiGPU: Successfully created class for {node_name}")
|
|
except Exception as e:
|
|
logging.error(f"MultiGPU: Error executing class definition for {node_name}: {str(e)}")
|
|
|
|
for node_name, node_class in node_classes.items():
|
|
NODE_CLASS_MAPPINGS[f"{node_name}MultiGPU"] = override_class(node_class)
|
|
logging.info(f"MultiGPU: Registered {node_name} with MultiGPU wrapper")
|
|
|
|
|
|
def check_module_exists(module_path):
|
|
"""Utility function to check if module exists"""
|
|
full_path = os.path.join("custom_nodes", module_path, "__init__.py")
|
|
logging.info(f"MultiGPU: Checking for module at {full_path}")
|
|
|
|
if not os.path.exists(full_path):
|
|
logging.info(f"MultiGPU: Module {module_path} not found - skipping")
|
|
return False
|
|
|
|
logging.info(f"MultiGPU: Found {module_path}, attempting to load")
|
|
return True
|
|
|
|
|
|
def register_module(module_path, target_nodes):
|
|
try:
|
|
# For core nodes, skip module loading and just register from the global mappings
|
|
if not module_path:
|
|
logging.info("MultiGPU: Starting core node registration")
|
|
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
|
|
for node in target_nodes:
|
|
if node in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS[f"{node}MultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS[node])
|
|
logging.info(f"MultiGPU: Registered core node {node}")
|
|
else:
|
|
logging.info(f"MultiGPU: Core node {node} not found - this shouldn't happen!")
|
|
return
|
|
|
|
# For custom nodes, try to load the module first
|
|
full_path = os.path.join("custom_nodes", module_path, "__init__.py")
|
|
logging.info(f"MultiGPU: Checking for module at {full_path}")
|
|
|
|
if not os.path.exists(full_path):
|
|
logging.info(f"MultiGPU: Module {module_path} not found - skipping")
|
|
return
|
|
|
|
logging.info(f"MultiGPU: Found {module_path}, attempting to load")
|
|
spec = importlib.util.spec_from_file_location(module_path, full_path)
|
|
module = importlib.util.module_from_spec(spec)
|
|
spec.loader.exec_module(module)
|
|
logging.info(f"MultiGPU: Executed {module_path} initialization")
|
|
|
|
# Use the module's local dictionary instead of the global one
|
|
local_map_name = "NODE_CLASS_MAPPINGS"
|
|
local_map = getattr(module, local_map_name, None)
|
|
if not local_map:
|
|
logging.info(f"MultiGPU: {module_path} has no '{local_map_name}' dictionary, skipping override.")
|
|
return
|
|
|
|
all_defined_nodes = list(local_map.keys())
|
|
logging.info(f"MultiGPU: {module_path} local dict keys: {all_defined_nodes}")
|
|
|
|
for node in target_nodes:
|
|
if node in local_map:
|
|
mgpu_class = override_class(local_map[node])
|
|
NODE_CLASS_MAPPINGS[f"{node}MultiGPU"] = mgpu_class
|
|
logging.info(f"MultiGPU: Successfully wrapped {node} from {module_path}")
|
|
else:
|
|
logging.info(f"MultiGPU: Node '{node}' not found in {module_path}'s local dictionary")
|
|
|
|
except Exception as e:
|
|
logging.info(f"MultiGPU: Error processing {module_path}: {str(e)}")
|
|
|
|
def register_LTXmodule(module_path, node_list):
|
|
"""Hard-coded registration for LTX Video nodes"""
|
|
global NODE_CLASS_MAPPINGS
|
|
|
|
if not check_module_exists(module_path):
|
|
return
|
|
|
|
class LTXVLoader:
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
return {
|
|
"required": {
|
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),
|
|
{"tooltip": "The name of the checkpoint (model) to load."}),
|
|
"dtype": (["bfloat16", "float32"], {"default": "bfloat16"})
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MODEL", "VAE")
|
|
RETURN_NAMES = ("model", "vae")
|
|
FUNCTION = "load"
|
|
CATEGORY = "lightricks/LTXV"
|
|
TITLE = "LTXV Loader"
|
|
OUTPUT_NODE = False
|
|
|
|
def load(self, ckpt_name, dtype):
|
|
from nodes import NODE_CLASS_MAPPINGS
|
|
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
|
|
return original_loader.load(ckpt_name, dtype)
|
|
|
|
ltx_nodes = {
|
|
"LTXVLoader": LTXVLoader
|
|
}
|
|
|
|
for node_name in node_list:
|
|
if node_name in ltx_nodes:
|
|
NODE_CLASS_MAPPINGS[f"{node_name}MultiGPU"] = override_class(ltx_nodes[node_name])
|
|
logging.info(f"MultiGPU: Registered hard-coded LTX node {node_name}")
|
|
|
|
def register_Florence2module(module_path, node_list):
|
|
"""Hard-coded registration for Florence2 nodes"""
|
|
global NODE_CLASS_MAPPINGS
|
|
|
|
if not check_module_exists(module_path):
|
|
return
|
|
|
|
class DownloadAndLoadFlorence2Model:
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
return {"required": {
|
|
"model": ([
|
|
'microsoft/Florence-2-base',
|
|
'microsoft/Florence-2-base-ft',
|
|
'microsoft/Florence-2-large',
|
|
'microsoft/Florence-2-large-ft',
|
|
'HuggingFaceM4/Florence-2-DocVQA',
|
|
'thwri/CogFlorence-2.1-Large',
|
|
'thwri/CogFlorence-2.2-Large',
|
|
'gokaygokay/Florence-2-SD3-Captioner',
|
|
'gokaygokay/Florence-2-Flux-Large',
|
|
'MiaoshouAI/Florence-2-base-PromptGen-v1.5',
|
|
'MiaoshouAI/Florence-2-large-PromptGen-v1.5',
|
|
'MiaoshouAI/Florence-2-base-PromptGen-v2.0',
|
|
'MiaoshouAI/Florence-2-large-PromptGen-v2.0'
|
|
], {"default": 'microsoft/Florence-2-base'}),
|
|
"precision": (['fp16','bf16','fp32'], {"default": 'fp16'}),
|
|
"attention": (['flash_attention_2', 'sdpa', 'eager'], {"default": 'sdpa'}),
|
|
},
|
|
"optional": {
|
|
"lora": ("PEFTLORA",),
|
|
}}
|
|
|
|
RETURN_TYPES = ("FL2MODEL",)
|
|
RETURN_NAMES = ("florence2_model",)
|
|
FUNCTION = "loadmodel"
|
|
CATEGORY = "Florence2"
|
|
|
|
def loadmodel(self, model, precision, attention, lora=None):
|
|
# Get original node instance
|
|
from nodes import NODE_CLASS_MAPPINGS
|
|
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2Model"]()
|
|
|
|
# Use original node to load model
|
|
return original_loader.loadmodel(model, precision, attention, lora)
|
|
|
|
class Florence2ModelLoader:
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
return {"required": {
|
|
"model": ([item.name for item in Path(folder_paths.models_dir, "LLM").iterdir() if item.is_dir()],
|
|
{"tooltip": "models are expected to be in Comfyui/models/LLM folder"}),
|
|
"precision": (['fp16','bf16','fp32'],),
|
|
"attention": (['flash_attention_2', 'sdpa', 'eager'], {"default": 'sdpa'}),
|
|
},
|
|
"optional": {
|
|
"lora": ("PEFTLORA",),
|
|
}}
|
|
|
|
RETURN_TYPES = ("FL2MODEL",)
|
|
RETURN_NAMES = ("florence2_model",)
|
|
FUNCTION = "loadmodel"
|
|
CATEGORY = "Florence2"
|
|
|
|
def loadmodel(self, model, precision, attention, lora=None):
|
|
# Get original node instance
|
|
from nodes import NODE_CLASS_MAPPINGS
|
|
original_loader = NODE_CLASS_MAPPINGS["Florence2ModelLoader"]()
|
|
|
|
# Use original node to load model
|
|
return original_loader.loadmodel(model, precision, attention, lora)
|
|
|
|
florence2_nodes = {
|
|
"Florence2ModelLoader": Florence2ModelLoader,
|
|
"DownloadAndLoadFlorence2Model": DownloadAndLoadFlorence2Model
|
|
}
|
|
|
|
for node_name in node_list:
|
|
if node_name in florence2_nodes:
|
|
NODE_CLASS_MAPPINGS[f"{node_name}MultiGPU"] = override_class(florence2_nodes[node_name])
|
|
logging.info(f"MultiGPU: Registered hard-coded Florence2 node {node_name}")
|
|
|
|
# Register desired nodes
|
|
register_module("", ["UNETLoader", "VAELoader", "CLIPLoader", "DualCLIPLoader", "TripleCLIPLoader", "CheckpointLoaderSimple", "ControlNetLoader"])
|
|
register_module("ComfyUI-GGUF", ["UnetLoaderGGUF","UnetLoaderGGUFAdvanced","CLIPLoaderGGUF","DualCLIPLoaderGGUF","TripleCLIPLoaderGGUF"])
|
|
#register_module("x-flux-comfyui", ["LoadFluxControlNet"])
|
|
#register_Florence2module("ComfyUI-Florence2", ["Florence2ModelLoader", "DownloadAndLoadFlorence2Model"])
|
|
#register_LTXmodule("ComfyUI-LTXVideo", ["LTXVLoader"])
|
|
#register_module("ComfyUI-MMAudio", ["MMAudioFeatureUtilsLoader","MMAudioModelLoader","MMAudioSampler"])
|
|
#register_module("ComfyUI_bitsandbytes_NF4", ["CheckpointLoaderNF4",])
|
|
|
|
|
|
#register_module_new("ComfyUI-GGUF", ["UnetLoaderGGUF","UnetLoaderGGUFAdvanced","CLIPLoaderGGUF","DualCLIPLoaderGGUF","TripleCLIPLoaderGGUF"])
|
|
#register_module_new("x-flux-comfyui", ["LoadFluxControlNet"])
|
|
#register_module_new("ComfyUI-Florence2", ["Florence2ModelLoader", "DownloadAndLoadFlorence2Model"])
|
|
register_module_new("ComfyUI-LTXVideo", ["LTXVLoader"])
|
|
#register_module_new("ComfyUI-MMAudio", ["MMAudioFeatureUtilsLoader","MMAudioModelLoader","MMAudioSampler"])
|
|
#register_module_new("ComfyUI_bitsandbytes_NF4", ["CheckpointLoaderNF4",])
|
|
|
|
|
|
logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
|