feat(multigpu): generalize LTX module registration for custom nodes

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.
This commit is contained in:
John Pollock
2024-12-31 15:51:43 -06:00
parent 3f57e7fe83
commit 1d12e2ee1f
+34 -53
View File
@@ -48,20 +48,13 @@ def override_class(cls):
NODE_CLASS_MAPPINGS = {}
def register_module_new(module_path, target_nodes):
"""
Locates the source code definition of specified nodes within a module,
extracting the block from the 'class' declaration up to and including the second 'def' line,
and logs the definition. This is a self-contained function.
Args:
module_path: The path to the custom node module's directory.
target_nodes: A list of node class names to process.
"""
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):
@@ -88,39 +81,35 @@ def register_module_new(module_path, target_nodes):
if start_index != -1 and second_def_index != -1:
definition_lines = lines[start_index:second_def_index]
node_definition_block = "".join(definition_lines)
break # Found the definition, move to the next node
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:
logging.info(f"MultiGPU: Definition for node '{node_name}' in module '{module_path}':\n{node_definition_block}")
# Add the wrapper class definition here
wrapper_code_lines = [
f" # Get original node instance",
f" from nodes import NODE_CLASS_MAPPINGS",
f" original_loader = NODE_CLASS_MAPPINGS[\"{node_name}\"]()"
]
node_definition_block += "\n".join(wrapper_code_lines)
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)}")
# Extract the function name and arguments from the last line of the definition block
last_line = definition_lines[-1].strip()
if last_line.startswith("def "):
method_declaration = last_line[len("def "):].rstrip(":")
method_name, args_str = method_declaration.split("(", 1)
# Remove "self, " if present
args_str = args_str.replace("self, ", "").replace("self,", "").strip()
# Construct the return line
return_line = f" return original_loader.{method_name}({args_str}"
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")
node_definition_block += "\n" + return_line
logging.info(f"MultiGPU: Concatenated lines for node '{node_name}':\n{node_definition_block}")
else:
logging.warning(f"MultiGPU: Could not parse the last line of definition for node '{node_name}'.")
else:
logging.info(f"MultiGPU: Could not retrieve definition for node '{node_name}' in module '{module_path}'.")
def check_module_exists(module_path):
"""Utility function to check if module exists"""
@@ -210,16 +199,8 @@ def register_LTXmodule(module_path, node_list):
OUTPUT_NODE = False
def load(self, ckpt_name, dtype):
# Get original node instance
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
# Use original node to load model and VAE
# model, vae = original_loader.load(ckpt_name, dtype)
# Return original objects
#return (model, vae)
return original_loader.load(ckpt_name, dtype)
ltx_nodes = {
@@ -316,19 +297,19 @@ def register_Florence2module(module_path, node_list):
# 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("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-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",])
#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())}")