From caee8716b67a4615023e462740aa63362dc54893 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Mon, 30 Dec 2024 07:54:57 -0600 Subject: [PATCH] Refactor - intermediate step of Implementing MultiGPU node registration and class definition retrieval for custom nodes --- __init__.py | 233 ++++++++++++++++++++++++++++++---------------------- 1 file changed, 136 insertions(+), 97 deletions(-) diff --git a/__init__.py b/__init__.py index 072f070..583c4a5 100644 --- a/__init__.py +++ b/__init__.py @@ -1,29 +1,103 @@ -import time +import os +import ast import copy import torch -import comfy.model_management -import os -import importlib.util import logging +from typing import Dict, Type -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 find_node_definition(module_path: str, node_name: str) -> Dict: + """ + Finds a specific node class definition by searching Python files in the given path. + """ + search_dir = os.path.join("custom_nodes", module_path) + if not os.path.exists(search_dir): + logging.info(f"MultiGPU: No custom_nodes directory {module_path}, skipping") + return None + + py_files = [] + for root, _, files in os.walk(search_dir): + for file in files: + if file.endswith('.py'): + py_files.append(os.path.join(root, file)) + + if not py_files: + return None + + try: + for file_path in py_files: + with open(file_path, 'r', encoding='utf-8') as f: + tree = ast.parse(f.read()) -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) + for node in ast.walk(tree): + if isinstance(node, ast.ClassDef) and node.name == node_name: + class_info = { + 'input_types': None, + 'return_types': None, + 'function': None, + 'category': None + } -comfy.model_management.get_torch_device = get_torch_device_patched + for item in node.body: + if isinstance(item, ast.FunctionDef) and item.name == 'INPUT_TYPES': + if any(d.id == 'classmethod' for d in item.decorator_list if isinstance(d, ast.Name)): + for stmt in item.body: + if isinstance(stmt, ast.Return): + try: + class_info['input_types'] = ast.literal_eval(stmt.value) + except: + pass + + elif isinstance(item, ast.Assign): + for target in item.targets: + if isinstance(target, ast.Name): + try: + if target.id == 'RETURN_TYPES': + class_info['return_types'] = ast.literal_eval(item.value) + elif target.id == 'FUNCTION': + class_info['function'] = ast.literal_eval(item.value) + elif target.id == 'CATEGORY': + class_info['category'] = ast.literal_eval(item.value) + except: + pass + + return class_info + + except Exception as e: + logging.error(f"MultiGPU: Error scanning for {node_name}: {str(e)}") + + return None + +def create_multigpu_node(node_name: str, class_info: Dict) -> Type: + """Creates a MultiGPU version of the node""" + class MultiGPUNode: + @classmethod + def INPUT_TYPES(cls): + inputs = copy.deepcopy(class_info['input_types']) + if inputs is None: + inputs = {"required": {}, "optional": {}} + elif "required" not in inputs: + inputs["required"] = {} + + devices = ["cpu"] + [f"cuda:{i}" for i in range(torch.cuda.device_count())] + inputs["required"]["device"] = (devices,) + return inputs + + RETURN_TYPES = class_info['return_types'] if class_info['return_types'] is not None else tuple() + FUNCTION = "override" + CATEGORY = "multigpu" + + def override(self, *args, device="cpu", **kwargs): + global current_device + current_device = device + return args if isinstance(args, tuple) else (args,) + + MultiGPUNode.__name__ = f"{node_name}MultiGPU" + return MultiGPUNode def override_class(cls): + """Creates a MultiGPU version of a node class that preserves original functionality.""" class NodeOverride(cls): @classmethod def INPUT_TYPES(s): @@ -43,96 +117,61 @@ def override_class(cls): return NodeOverride -def register_module(module_path, target_nodes): +def register_core_nodes(target_nodes: list): + """Register MultiGPU versions of core ComfyUI nodes""" try: - # For core nodes, skip module loading and just register from 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 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") + logging.info("MultiGPU: Processing core 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}") + logging.info(f"MultiGPU: Registered {node}MultiGPU") else: - logging.info(f"MultiGPU: Node {node} from {module_path} not found in global mappings") + logging.info(f"MultiGPU: Core node {node} not found") + except Exception as e: + logging.error(f"MultiGPU: Error processing core nodes: {str(e)}") + +def register_module(module_path: str, target_nodes: list): + """Register MultiGPU versions of custom nodes""" + try: + # Handle custom nodes + logging.info(f"MultiGPU: Processing module {module_path}") + search_dir = os.path.join("custom_nodes", module_path) + if not os.path.exists(search_dir): + logging.info(f"MultiGPU: Module directory {module_path} not found, skipping") + return + + # List all Python files in the module once + py_files = [] + for root, _, files in os.walk(search_dir): + for file in files: + if file.endswith('.py'): + py_files.append(os.path.basename(file)) + + if py_files: + logging.info(f"MultiGPU: Searching in {module_path}: {', '.join(sorted(py_files))}") + + for node in target_nodes: + class_info = find_node_definition(module_path, node) + if class_info: + NODE_CLASS_MAPPINGS[f"{node}MultiGPU"] = create_multigpu_node(node, class_info) + logging.info(f"MultiGPU: Registered {node}MultiGPU") except Exception as e: - logging.info(f"MultiGPU: Error processing {module_path}: {str(e)}") + logging.error(f"MultiGPU: Error in {module_path}: {str(e)}") +# Initialize NODE_CLASS_MAPPINGS = {} +current_device = None -# Let's test just one new module at a time, starting with GGUF -logging.info("MultiGPU: Starting Core ComfyUI registration") -register_module("", [ - "UNETLoader", - "VAELoader", - "CLIPLoader", - "DualCLIPLoader", - "TripleCLIPLoader", - "CheckpointLoaderSimple", - "ControlNetLoader" -]) +# Register all modules +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_module("ComfyUI-Florence2", ["Florence2ModelLoader", "DownloadAndLoadFlorence2Model"]) +register_module("ComfyUI-LTXVideo", ["LTXVLoader"]) +register_module("ComfyUI-MMAudio", [ "MMAudioFeatureUtilsLoader", "MMAudioModelLoader", "MMAudioSampler"]) +register_module("ComfyUI_bitsandbytes_NF4", ["CheckpointLoaderNF4"]) -logging.info("MultiGPU: Starting GGUF registration") -register_module("ComfyUI-GGUF", [ - "UnetLoaderGGUF", - "UnetLoaderGGUFAdvanced", - "CLIPLoaderGGUF", - "DualCLIPLoaderGGUF", - "TripleCLIPLoaderGGUF" -]) -logging.info("MultiGPU: Starting X-Flux ControlNet registration") -register_module("x-flux-comfyui", [ - "LoadFluxControlNet" -]) - -logging.info("MultiGPU: Starting Florence2 registration") -register_module("ComfyUI-Florence2", [ - "Florence2ModelLoader", - "DownloadAndLoadFlorence2Model" -]) - -logging.info("MultiGPU: Starting LTXVideo registration") -register_module("ComfyUI-LTXVideo", [ - "LTXVLoader" -]) - -logging.info("MultiGPU: Starting MMAudio registration") -register_module("ComfyUI-MMAudio", [ - "MMAudioFeatureUtilsLoader", - "MMAudioModelLoader", - "MMAudioSampler" -]) - -logging.info("MultiGPU: Starting NF4 registration") -register_module("ComfyUI_bnb_nf4_fp4_Loaders", [ - "CheckpointLoaderNF4", - "UNETLoaderNF4" -]) - - -logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}") \ No newline at end of file +logging.info("MultiGPU: Registration complete") +logging.info(f"MultiGPU: Registered nodes: {', '.join(sorted(NODE_CLASS_MAPPINGS.keys()))}") \ No newline at end of file