From 4ac8d3327054d8f8b6f911f69ef5f488b1dc5cef Mon Sep 17 00:00:00 2001 From: John Pollock Date: Mon, 30 Dec 2024 11:17:31 -0600 Subject: [PATCH] fix: Use local module dictionary instead of global for custom node registration - Added a local_map_name ("NODE_CLASS_MAPPINGS") and retrieve it from the module immediately after loading. - If the local dictionary exists, wrap the target nodes from there, rather than relying on the global dictionary. - Removed references to GLOBAL_NODE_CLASS_MAPPINGS for custom nodes and replaced them with the local module mapping lookup. --- __init__.py | 40 +++++++++++++++++++++++++--------------- 1 file changed, 25 insertions(+), 15 deletions(-) diff --git a/__init__.py b/__init__.py index 25e6da9..0f72ebf 100644 --- a/__init__.py +++ b/__init__.py @@ -43,9 +43,11 @@ def override_class(cls): return NodeOverride +NODE_CLASS_MAPPINGS = {} + def register_module(module_path, target_nodes): try: - # For core nodes, skip module loading and just register from global mappings + # 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 @@ -57,34 +59,42 @@ def register_module(module_path, target_nodes): logging.info(f"MultiGPU: Core node {node} not found - this shouldn't happen!") return - # For custom nodes, try to load module first + # 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") - - from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS - logging.info(f"MultiGPU: Looking for {module_path} nodes in global mappings") + + # 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 GLOBAL_NODE_CLASS_MAPPINGS: - NODE_CLASS_MAPPINGS[f"{node}MultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS[node]) - logging.info(f"MultiGPU: Successfully wrapped {node}") + 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} from {module_path} not found in global mappings") - + 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)}") -NODE_CLASS_MAPPINGS = {} - +# 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"]) @@ -93,4 +103,4 @@ register_module("ComfyUI-LTXVideo", ["LTXVLoader"]) register_module("ComfyUI-MMAudio", ["MMAudioFeatureUtilsLoader","MMAudioModelLoader","MMAudioSampler"]) register_module("ComfyUI_bitsandbytes_NF4", ["CheckpointLoaderNF4",]) -logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}") \ No newline at end of file +logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")