From 7af256ab6ea86b105777631fce116c4174547362 Mon Sep 17 00:00:00 2001 From: Your Name Date: Fri, 6 Mar 2026 01:35:50 +0000 Subject: [PATCH] fix: harden startup compat and restore low-risk MultiGPU helpers --- __init__.py | 179 +++++++++++++++++++++++++++++++++--------------- device_utils.py | 1 + nodes.py | 16 +++++ wanvideo.py | 9 ++- 4 files changed, 148 insertions(+), 57 deletions(-) diff --git a/__init__.py b/__init__.py index 2f5efd4..b2c4148 100644 --- a/__init__.py +++ b/__init__.py @@ -4,6 +4,7 @@ import weakref import os import copy import json +import importlib from datetime import datetime from pathlib import Path import folder_paths @@ -136,15 +137,36 @@ def mgpu_mm_log_method(self, msg): ) logger.mgpu_mm_log = mgpu_mm_log_method.__get__(logger, type(logger)) +def _normalize_module_name(module_name): + """Normalize a custom node directory name for tolerant matching.""" + return "".join(char for char in os.path.basename(module_name).lower() if char.isalnum()) + def check_module_exists(module_path): """Check if a custom node module exists in ComfyUI custom_nodes directory.""" - full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path) - logger.debug(f"[MultiGPU] Checking for module at {full_path}") - if not os.path.exists(full_path): - logger.debug(f"[MultiGPU] Module {module_path} not found - skipping") - return False - logger.debug(f"[MultiGPU] Found {module_path}, creating compatible MultiGPU nodes") - return True + custom_nodes_paths = folder_paths.get_folder_paths("custom_nodes") + normalized_module_path = _normalize_module_name(module_path) + + for custom_nodes_path in custom_nodes_paths: + full_path = os.path.join(custom_nodes_path, module_path) + logger.debug(f"[MultiGPU] Checking for module at {full_path}") + if os.path.isdir(full_path): + logger.debug(f"[MultiGPU] Found exact module match for {module_path} at {full_path}") + return True + + for custom_nodes_path in custom_nodes_paths: + try: + with os.scandir(custom_nodes_path) as entries: + for entry in entries: + if not entry.is_dir(): + continue + if _normalize_module_name(entry.name) == normalized_module_path: + logger.debug(f"[MultiGPU] Found normalized module match for {module_path} at {entry.path}") + return True + except OSError: + continue + + logger.debug(f"[MultiGPU] Module {module_path} not found - skipping") + return False current_device = mm.get_torch_device() current_text_encoder_device = mm.text_encoder_device() @@ -216,6 +238,44 @@ def unet_offload_device_patched(): logger.debug(f"[MultiGPU Core Patching] unet_offload_device_patched returning device: {device} (current_unet_offload_device={current_unet_offload_device})") return device +def _patch_comfy_kitchen_dlpack_device_guard(): + """Guard comfy_kitchen DLPack export by switching to the tensor's CUDA device.""" + try: + comfy_kitchen_cuda = importlib.import_module("comfy_kitchen.backends.cuda") + except ImportError: + logger.debug("[MultiGPU] comfy_kitchen not found - skipping CUDA DLPack compat patch") + return False + + wrap_for_dlpack = getattr(comfy_kitchen_cuda, "_wrap_for_dlpack", None) + if wrap_for_dlpack is None: + logger.debug("[MultiGPU] comfy_kitchen.backends.cuda._wrap_for_dlpack not found - skipping compat patch") + return False + + if getattr(wrap_for_dlpack, "_multigpu_cuda_device_guard", False): + return True + + def wrap_for_dlpack_with_device_guard(*args, **kwargs): + tensor = args[0] if args else kwargs.get("tensor") + previous_device_index = None + switched_device = False + + if isinstance(tensor, torch.Tensor) and tensor.is_cuda and tensor.device.index is not None: + previous_device_index = torch.cuda.current_device() + if previous_device_index != tensor.device.index: + torch.cuda.set_device(tensor.device.index) + switched_device = True + + try: + return wrap_for_dlpack(*args, **kwargs) + finally: + if switched_device and previous_device_index is not None: + torch.cuda.set_device(previous_device_index) + + wrap_for_dlpack_with_device_guard._multigpu_cuda_device_guard = True + comfy_kitchen_cuda._wrap_for_dlpack = wrap_for_dlpack_with_device_guard + logger.info("[MultiGPU] Applied comfy_kitchen CUDA DLPack device guard patch") + return True + logger.info(f"[MultiGPU Core Patching] Patching mm.get_torch_device, mm.text_encoder_device, mm.unet_offload_device") logger.info(f"[MultiGPU DEBUG] Initial current_device: {current_device}") logger.info(f"[MultiGPU DEBUG] Initial current_text_encoder_device: {current_text_encoder_device}") @@ -224,8 +284,10 @@ logger.info(f"[MultiGPU DEBUG] Initial current_unet_offload_device: {current_une mm.get_torch_device = get_torch_device_patched mm.text_encoder_device = text_encoder_device_patched mm.unet_offload_device = unet_offload_device_patched +_patch_comfy_kitchen_dlpack_device_guard() from .nodes import ( + DeviceSelectorMultiGPU, UnetLoaderGGUF, UnetLoaderGGUFAdvanced, CLIPLoaderGGUF, @@ -246,29 +308,6 @@ from .nodes import ( UNetLoaderLP, ) -from .wanvideo import ( - LoadWanVideoT5TextEncoder, - WanVideoTextEncode, - WanVideoTextEncodeCached, - WanVideoTextEncodeSingle, - WanVideoVAELoader, - WanVideoTinyVAELoader, - WanVideoBlockSwap, - WanVideoImageToVideoEncode, - WanVideoDecode, - WanVideoModelLoader, - WanVideoSampler, - WanVideoVACEEncode, - WanVideoEncode, - LoadWanVideoClipTextEncoder, - WanVideoClipVisionEncode, - WanVideoControlnetLoader, - FantasyTalkingModelLoader, - Wav2VecModelLoader, - WanVideoUni3C_ControlnetLoader, - DownloadAndLoadWav2VecModel, -) - from .wrappers import ( override_class, override_class_offload, @@ -294,9 +333,57 @@ from .checkpoint_multigpu import ( CheckpointLoaderAdvancedDisTorch2MultiGPU ) +def _load_wanvideo_nodes(): + from .wanvideo import ( + LoadWanVideoT5TextEncoder, + WanVideoTextEncode, + WanVideoTextEncodeCached, + WanVideoTextEncodeSingle, + WanVideoVAELoader, + WanVideoTinyVAELoader, + WanVideoBlockSwap, + WanVideoImageToVideoEncode, + WanVideoDecode, + WanVideoModelLoader, + WanVideoSampler, + WanVideoVACEEncode, + WanVideoEncode, + LoadWanVideoClipTextEncoder, + WanVideoClipVisionEncode, + WanVideoControlnetLoader, + FantasyTalkingModelLoader, + Wav2VecModelLoader, + WanVideoUni3C_ControlnetLoader, + DownloadAndLoadWav2VecModel, + ) + + return { + "LoadWanVideoT5TextEncoderMultiGPU": LoadWanVideoT5TextEncoder, + "WanVideoTextEncodeMultiGPU": WanVideoTextEncode, + "WanVideoTextEncodeCachedMultiGPU": WanVideoTextEncodeCached, + "WanVideoTextEncodeSingleMultiGPU": WanVideoTextEncodeSingle, + "WanVideoVAELoaderMultiGPU": WanVideoVAELoader, + "WanVideoTinyVAELoaderMultiGPU": WanVideoTinyVAELoader, + "WanVideoBlockSwapMultiGPU": WanVideoBlockSwap, + "WanVideoImageToVideoEncodeMultiGPU": WanVideoImageToVideoEncode, + "WanVideoDecodeMultiGPU": WanVideoDecode, + "WanVideoModelLoaderMultiGPU": WanVideoModelLoader, + "WanVideoSamplerMultiGPU": WanVideoSampler, + "WanVideoVACEEncodeMultiGPU": WanVideoVACEEncode, + "WanVideoEncodeMultiGPU": WanVideoEncode, + "LoadWanVideoClipTextEncoderMultiGPU": LoadWanVideoClipTextEncoder, + "WanVideoClipVisionEncodeMultiGPU": WanVideoClipVisionEncode, + "WanVideoControlnetLoaderMultiGPU": WanVideoControlnetLoader, + "FantasyTalkingModelLoaderMultiGPU": FantasyTalkingModelLoader, + "Wav2VecModelLoaderMultiGPU": Wav2VecModelLoader, + "WanVideoUni3C_ControlnetLoaderMultiGPU": WanVideoUni3C_ControlnetLoader, + "DownloadAndLoadWav2VecModelMultiGPU": DownloadAndLoadWav2VecModel, + } + NODE_CLASS_MAPPINGS = { "CheckpointLoaderAdvancedMultiGPU": CheckpointLoaderAdvancedMultiGPU, "CheckpointLoaderAdvancedDisTorch2MultiGPU": CheckpointLoaderAdvancedDisTorch2MultiGPU, + "DeviceSelectorMultiGPU": DeviceSelectorMultiGPU, "UNetLoaderLP": UNetLoaderLP, } @@ -342,8 +429,14 @@ def register_and_count(module_names, node_map): count = 0 if found: + try: + resolved_node_map = node_map() if callable(node_map) else node_map + except Exception as exc: + logger.warning(f"[MultiGPU] Failed to register nodes for {module_names[0]}: {exc}") + resolved_node_map = {} + initial_len = len(NODE_CLASS_MAPPINGS) - for key, value in node_map.items(): + for key, value in resolved_node_map.items(): NODE_CLASS_MAPPINGS[key] = value count = len(NODE_CLASS_MAPPINGS) - initial_len @@ -401,29 +494,7 @@ pulid_nodes = { } register_and_count(["PuLID_ComfyUI", "pulid_comfyui"], pulid_nodes) -wanvideo_nodes = { - "LoadWanVideoT5TextEncoderMultiGPU": LoadWanVideoT5TextEncoder, - "WanVideoTextEncodeMultiGPU": WanVideoTextEncode, - "WanVideoTextEncodeCachedMultiGPU": WanVideoTextEncodeCached, - "WanVideoTextEncodeSingleMultiGPU": WanVideoTextEncodeSingle, - "WanVideoVAELoaderMultiGPU": WanVideoVAELoader, - "WanVideoTinyVAELoaderMultiGPU": WanVideoTinyVAELoader, - "WanVideoBlockSwapMultiGPU": WanVideoBlockSwap, - "WanVideoImageToVideoEncodeMultiGPU": WanVideoImageToVideoEncode, - "WanVideoDecodeMultiGPU": WanVideoDecode, - "WanVideoModelLoaderMultiGPU": WanVideoModelLoader, - "WanVideoSamplerMultiGPU": WanVideoSampler, - "WanVideoVACEEncodeMultiGPU": WanVideoVACEEncode, - "WanVideoEncodeMultiGPU": WanVideoEncode, - "LoadWanVideoClipTextEncoderMultiGPU": LoadWanVideoClipTextEncoder, - "WanVideoClipVisionEncodeMultiGPU": WanVideoClipVisionEncode, - "WanVideoControlnetLoaderMultiGPU": WanVideoControlnetLoader, - "FantasyTalkingModelLoaderMultiGPU": FantasyTalkingModelLoader, - "Wav2VecModelLoaderMultiGPU": Wav2VecModelLoader, - "WanVideoUni3C_ControlnetLoaderMultiGPU": WanVideoUni3C_ControlnetLoader, - "DownloadAndLoadWav2VecModelMultiGPU": DownloadAndLoadWav2VecModel, -} -register_and_count(["ComfyUI-WanVideoWrapper", "comfyui-wanvideowrapper"], wanvideo_nodes) +register_and_count(["ComfyUI-WanVideoWrapper", "comfyui-wanvideowrapper"], _load_wanvideo_nodes) for item in registration_data: logger.info(fmt_reg.format(item['name'], item['found'], str(item['count']))) diff --git a/device_utils.py b/device_utils.py index 6647572..6b7b7ae 100644 --- a/device_utils.py +++ b/device_utils.py @@ -175,6 +175,7 @@ def soft_empty_cache_multigpu(): logger.mgpu_mm_log(f"Clearing CUDA cache on {device_str} (idx={device_idx})") multigpu_memory_log("general", f"pre-empty:{device_str}") with torch.cuda.device(device_idx): + torch.cuda.synchronize() torch.cuda.empty_cache() if hasattr(torch.cuda, "ipc_collect"): torch.cuda.ipc_collect() diff --git a/nodes.py b/nodes.py index 4a75b78..c2d3787 100644 --- a/nodes.py +++ b/nodes.py @@ -5,6 +5,22 @@ from nodes import NODE_CLASS_MAPPINGS from .device_utils import get_device_list from .model_management_mgpu import force_full_system_cleanup +class DeviceSelectorMultiGPU: + @classmethod + def INPUT_TYPES(s): + devices = get_device_list() + return {"required": {"device": (devices,)}} + + RETURN_TYPES = ("MULTIGPUDEVICE",) + RETURN_NAMES = ("device",) + FUNCTION = "select_device" + CATEGORY = "multigpu" + TITLE = "Device Selector (MultiGPU)" + + def select_device(self, device): + """Return the selected device label without side effects.""" + return (device,) + class UnetLoaderGGUF: @classmethod def INPUT_TYPES(s): diff --git a/wanvideo.py b/wanvideo.py index 4ebecf9..1628e26 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -11,9 +11,7 @@ from .model_management_mgpu import multigpu_memory_log from comfy.utils import load_torch_file, ProgressBar import gc import numpy as np -from accelerate import init_empty_weights import os -import importlib.util logger = logging.getLogger("MultiGPU") @@ -64,6 +62,7 @@ class WanVideoModelLoader: "lora": ("WANVIDLORA", {"default": None}), "vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}), "extra_model": ("VACEPATH", {"default": None, "tooltip": "Extra model to add to the main model, ie. VACE or MTV Crafter"}), + "vace_model": ("VACEPATH", {"default": None, "tooltip": "Backward-compatible alias for extra_model"}), "fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}), "multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}), "fantasyportrait_model": ("FANTASYPORTRAITMODEL", {"default": None, "tooltip": "FantasyPortrait model"}), @@ -83,6 +82,10 @@ class WanVideoModelLoader: loader_module = inspect.getmodule(original_loader) original_module_device = loader_module.device + vace_model = kwargs.pop("vace_model", None) + if kwargs.get("extra_model") is None and vace_model is not None: + kwargs["extra_model"] = vace_model + set_current_device(compute_device) compute_device_to_be_patched = mm.get_torch_device() @@ -863,4 +866,4 @@ class WanVideoUni3C_ControlnetLoader: set_current_device(device) original_loader = NODE_CLASS_MAPPINGS["WanVideoUni3C_ControlnetLoader"]() - return original_loader.loadmodel(model, base_precision, load_device, quantization, attention_mode, compile_args) \ No newline at end of file + return original_loader.loadmodel(model, base_precision, load_device, quantization, attention_mode, compile_args)