diff --git a/__init__.py b/__init__.py index 268b90a..6c42fbd 100644 --- a/__init__.py +++ b/__init__.py @@ -28,8 +28,8 @@ from .nodes import ( MMAudioModelLoader, MMAudioFeatureUtilsLoader, MMAudioSampler, PulidModelLoader, PulidInsightFaceLoader, PulidEvaClipLoader, HyVideoModelLoader, HyVideoVAELoader, DownloadAndLoadHyVideoTextEncoder, - WanVideoModelLoader, WanVideoVAELoader, LoadWanVideoT5TextEncoder, LoadWanVideoClipTextEncoder, - WanVideoTextEncode, WanVideoBlockSwap, WanVideoModelLoader_TWO + WanVideoModelLoader, WanVideoModelLoader_2, WanVideoVAELoader, LoadWanVideoT5TextEncoder, LoadWanVideoClipTextEncoder, + WanVideoTextEncode, WanVideoBlockSwap, WanVideoSampler ) current_device = mm.get_torch_device() @@ -761,12 +761,13 @@ if check_module_exists("ComfyUI-HunyuanVideoWrapper") or check_module_exists("co if check_module_exists("ComfyUI-WanVideoWrapper") or check_module_exists("comfyui-wanvideowrapper"): # WanVideo uses custom implementation, not the standard override NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU"] = WanVideoModelLoader - NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU_TWO"] = WanVideoModelLoader_TWO + NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU_2"] = WanVideoModelLoader_2 NODE_CLASS_MAPPINGS["WanVideoVAELoaderMultiGPU"] = WanVideoVAELoader NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoderMultiGPU"] = LoadWanVideoT5TextEncoder NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoderMultiGPU"] = LoadWanVideoClipTextEncoder NODE_CLASS_MAPPINGS["WanVideoTextEncodeMultiGPU"] = WanVideoTextEncode NODE_CLASS_MAPPINGS["WanVideoBlockSwapMultiGPU"] = WanVideoBlockSwap + NODE_CLASS_MAPPINGS["WanVideoSamplerMultiGPU"] = WanVideoSampler logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}") diff --git a/nodes.py b/nodes.py index ff4a9c7..c1f94fa 100644 --- a/nodes.py +++ b/nodes.py @@ -882,125 +882,62 @@ class LoadWanVideoClipTextEncoder: return original_loader.loadmodel(model_name, precision, load_device) -class WanVideoModelLoader_TWO: - """Second instance of WanVideoModelLoader for multi-model workflows to avoid race conditions""" + +class WanVideoModelLoader_2: + """Second instance for multi-model workflows to maintain separate device patches""" @classmethod def INPUT_TYPES(s): - # Exact same inputs as WanVideoModelLoader - from . import get_device_list - devices = get_device_list() - - return { - "required": { - "model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), - {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' folder",}), - "base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}), - "quantization": ( - ["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp8_e4m3fn_fast_no_ffn", "fp8_e4m3fn_scaled", "fp8_e5m2_scaled"], - {"default": "disabled", "tooltip": "optional quantization method"} - ), - "device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0], "tooltip": "Device to load the model to"}), - }, - "optional": { - "attention_mode": ([ - "sdpa", - "flash_attn_2", - "flash_attn_3", - "sageattn", - "sageattn_3", - "flex_attention", - "radial_sage_attention", - ], {"default": "sdpa"}), - "compile_args": ("WANCOMPILEARGS", ), - "block_swap_args": ("BLOCKSWAPARGS", ), - "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"}), - "vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}), - "fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}), - "multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}), - } - } - RETURN_TYPES = ("WANVIDEOMODEL",) - RETURN_NAMES = ("model", ) + # Delegate to the primary loader + return WanVideoModelLoader.INPUT_TYPES() + + RETURN_TYPES = WanVideoModelLoader.RETURN_TYPES + RETURN_NAMES = WanVideoModelLoader.RETURN_NAMES FUNCTION = "loadmodel" CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Second model loader for multi-model workflows - avoids race conditions when loading multiple models" - + DESCRIPTION = "Second model loader instance for workflows using multiple models on different devices" + def loadmodel(self, model, base_precision, device, quantization, - compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None): - # Just call the first loader's implementation directly - import logging - import comfy.model_management as mm - import torch - - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] ========== CUSTOM IMPLEMENTATION ==========") - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] User selected device: {device}") - - # Convert device string to torch device - selected_device = torch.device(device) - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Torch device: {selected_device}") - - # Determine load_device parameter for original loader - # If user selected CPU, use "offload_device", otherwise use "main_device" - load_device = "offload_device" if device == "cpu" else "main_device" - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Mapped to load_device: {load_device}") - + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, + vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None): + # Just use the first loader's implementation + loader = WanVideoModelLoader() + return loader.loadmodel(model, base_precision, device, quantization, + compile_args, attention_mode, block_swap_args, lora, + vram_management_args, vace_model, fantasytalking_model, multitalk_model) + + +class WanVideoSampler: + """Wrapper that ensures correct device patching before sampling""" + @classmethod + def INPUT_TYPES(s): + # Get original sampler's inputs from nodes import NODE_CLASS_MAPPINGS - original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]() - - # Patch BOTH WanVideo modules with the selected device + original_types = NODE_CLASS_MAPPINGS["WanVideoSampler"].INPUT_TYPES() + return original_types + + RETURN_TYPES = ("LATENT", "LATENT",) + RETURN_NAMES = ("samples", "denoised_samples",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "MultiGPU-aware sampler that ensures correct device for each model" + + def process(self, model, **kwargs): + import logging import sys - import inspect - loader_module = inspect.getmodule(original_loader) - if loader_module: - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Patching WanVideo modules to use {selected_device}") - - # Save original devices - original_device = getattr(loader_module, 'device', None) - original_offload = getattr(loader_module, 'offload_device', None) - - # Check if there's a model offload device override (from block swap config) - model_offload_override = getattr(loader_module, '_model_offload_device_override', None) - - # Patch nodes_model_loading.py module - setattr(loader_module, 'device', selected_device) - if model_offload_override: - # Use the model offload override for offload_device - setattr(loader_module, 'offload_device', model_offload_override) - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Using model offload override: {model_offload_override}") - elif device == "cpu": - setattr(loader_module, 'offload_device', selected_device) - - # Patch nodes.py module as well - nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes') - if nodes_module_name in sys.modules: - nodes_module = sys.modules[nodes_module_name] - setattr(nodes_module, 'device', selected_device) - - # Check for model offload override in nodes module too - nodes_model_offload_override = getattr(nodes_module, '_model_offload_device_override', None) - if nodes_model_offload_override: - setattr(nodes_module, 'offload_device', nodes_model_offload_override) - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Using model offload override for nodes.py: {nodes_model_offload_override}") - elif device == "cpu": - setattr(nodes_module, 'offload_device', selected_device) - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Both modules patched successfully") - - # Call original loader with our patches in place - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Calling original loader with patched device") - result = original_loader.loadmodel(model, base_precision, load_device, quantization, - compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model) - - # Leave patches in place for subsequent operations - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] Model loaded on {selected_device}") - logging.info(f"[MultiGPU WanVideoModelLoader_TWO] ========== COMPLETE ==========") - - return result - else: - logging.error(f"[MultiGPU WanVideoModelLoader_TWO] Could not patch modules, falling back") - return original_loader.loadmodel(model, base_precision, load_device, quantization, - compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model) + # Get the model's device and update WanVideo modules to match + model_device = model.load_device + logging.info(f"[MultiGPU WanVideoSampler] Model device: {model_device}") + + # Update the device variable in WanVideo modules + for module_name in sys.modules.keys(): + if 'WanVideoWrapper' in module_name and hasattr(sys.modules[module_name], 'device'): + sys.modules[module_name].device = model_device + + # Call original sampler + from nodes import NODE_CLASS_MAPPINGS + original_sampler = NODE_CLASS_MAPPINGS["WanVideoSampler"]() + return original_sampler.process(model, **kwargs) class WanVideoBlockSwap: