fix: harden startup compat and restore low-risk MultiGPU helpers

This commit is contained in:
Your Name
2026-03-06 01:35:50 +00:00
parent ac3df4ed70
commit 7af256ab6e
4 changed files with 148 additions and 57 deletions
+125 -54
View File
@@ -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'])))
+1
View File
@@ -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()
+16
View File
@@ -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):
+6 -3
View File
@@ -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)
return original_loader.loadmodel(model, base_precision, load_device, quantization, attention_mode, compile_args)