feat: add WanVideoBlockSwap class and integrate into MultiGPU node mappings
This commit is contained in:
@@ -215,6 +215,7 @@ from .wanvideo import (
|
||||
WanVideoTextEncode,
|
||||
WanVideoVAELoader,
|
||||
WanVideoTinyVAELoader,
|
||||
WanVideoBlockSwap,
|
||||
WanVideoImageToVideoEncode,
|
||||
WanVideoDecode,
|
||||
WanVideoModelLoader,
|
||||
@@ -359,6 +360,7 @@ wanvideo_nodes = {
|
||||
"WanVideoTextEncodeMultiGPU": WanVideoTextEncode,
|
||||
"WanVideoVAELoaderMultiGPU": WanVideoVAELoader,
|
||||
"WanVideoTinyVAELoaderMultiGPU": WanVideoTinyVAELoader,
|
||||
"WanVideoBlockSwapMultiGPU": WanVideoBlockSwap,
|
||||
"WanVideoImageToVideoEncodeMultiGPU": WanVideoImageToVideoEncode,
|
||||
"WanVideoDecodeMultiGPU": WanVideoDecode,
|
||||
"WanVideoModelLoaderMultiGPU": WanVideoModelLoader,
|
||||
|
||||
+178
-61
@@ -11,6 +11,7 @@ import logging
|
||||
import torch
|
||||
import sys
|
||||
import inspect
|
||||
import copy
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
from nodes import NODE_CLASS_MAPPINGS
|
||||
@@ -214,6 +215,35 @@ class WanVideoTinyVAELoader:
|
||||
# Return both the VAE model AND the selected device for device propagation
|
||||
return vae_model, load_device
|
||||
|
||||
|
||||
class WanVideoBlockSwap:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
base_inputs = copy.deepcopy(NODE_CLASS_MAPPINGS["WanVideoBlockSwap"].INPUT_TYPES())
|
||||
devices = get_device_list()
|
||||
default_device = "cpu" if "cpu" in devices else devices[0]
|
||||
base_inputs.setdefault("optional", {})
|
||||
base_inputs["optional"]["swap_device"] = (
|
||||
devices,
|
||||
{
|
||||
"default": default_device,
|
||||
"tooltip": "Device that receives swapped transformer blocks",
|
||||
},
|
||||
)
|
||||
return base_inputs
|
||||
|
||||
RETURN_TYPES = ("BLOCKSWAPARGS",)
|
||||
RETURN_NAMES = ("block_swap_args",)
|
||||
FUNCTION = "setargs"
|
||||
CATEGORY = "multigpu/WanVideoWrapper"
|
||||
DESCRIPTION = "Extends Wan block swap with explicit device selection"
|
||||
|
||||
def setargs(self, swap_device=None, **kwargs):
|
||||
block_swap_config = dict(kwargs)
|
||||
if swap_device is not None:
|
||||
block_swap_config["swap_device"] = str(swap_device)
|
||||
return (block_swap_config,)
|
||||
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -540,61 +570,117 @@ class WanVideoModelLoader:
|
||||
fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None,
|
||||
rms_norm_function="default"):
|
||||
from . import set_current_device
|
||||
logger.info(
|
||||
f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] User selected device: {compute_device}"
|
||||
)
|
||||
|
||||
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] User selected device: {compute_device}")
|
||||
|
||||
selected_device = torch.device(compute_device)
|
||||
|
||||
# Set the global device context for any downstream operations
|
||||
set_current_device(selected_device)
|
||||
|
||||
# Find the original loader and its module
|
||||
|
||||
normalized_block_swap = None
|
||||
swap_device_override = None
|
||||
if block_swap_args is not None:
|
||||
normalized_block_swap = dict(block_swap_args)
|
||||
swap_selection = normalized_block_swap.pop("swap_device", None)
|
||||
if swap_selection is None:
|
||||
swap_selection = normalized_block_swap.get("resolved_swap_device")
|
||||
if swap_selection is None:
|
||||
swap_selection = "cpu"
|
||||
try:
|
||||
swap_device_override = torch.device(str(swap_selection))
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Invalid swap_device '%s', falling back to CPU",
|
||||
swap_selection,
|
||||
)
|
||||
swap_device_override = torch.device("cpu")
|
||||
normalized_block_swap["resolved_swap_device"] = str(swap_device_override)
|
||||
|
||||
original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]()
|
||||
loader_module = inspect.getmodule(original_loader)
|
||||
|
||||
if loader_module:
|
||||
logger.debug(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Patching '{loader_module.__name__}' to use device: {selected_device}")
|
||||
|
||||
# Store original values to restore later if needed, though it's less critical in this workflow
|
||||
original_module_device = getattr(loader_module, 'device', None)
|
||||
|
||||
# Overwrite the module-level 'device' variable. This is the key to the fix.
|
||||
setattr(loader_module, 'device', selected_device)
|
||||
|
||||
# Also patch the offload device if the user chose CPU
|
||||
if compute_device == "cpu":
|
||||
setattr(loader_module, 'offload_device', selected_device)
|
||||
|
||||
logger.debug("[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Device patching complete. Calling original loader...")
|
||||
|
||||
# Call the original loader function with all the arguments it expects
|
||||
if not loader_module:
|
||||
logger.error(
|
||||
"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Could not resolve loader module; invoking original implementation without patches."
|
||||
)
|
||||
result = original_loader.loadmodel(
|
||||
model, base_precision, load_device, quantization,
|
||||
compile_args, attention_mode, block_swap_args, lora,
|
||||
vram_management_args, extra_model=extra_model, vace_model=vace_model,
|
||||
fantasytalking_model=fantasytalking_model, multitalk_model=multitalk_model,
|
||||
fantasyportrait_model=fantasyportrait_model, rms_norm_function=rms_norm_function
|
||||
model,
|
||||
base_precision,
|
||||
load_device,
|
||||
quantization,
|
||||
compile_args,
|
||||
attention_mode,
|
||||
normalized_block_swap if normalized_block_swap is not None else block_swap_args,
|
||||
lora,
|
||||
vram_management_args,
|
||||
extra_model=extra_model,
|
||||
vace_model=vace_model,
|
||||
fantasytalking_model=fantasytalking_model,
|
||||
multitalk_model=multitalk_model,
|
||||
fantasyportrait_model=fantasyportrait_model,
|
||||
rms_norm_function=rms_norm_function,
|
||||
)
|
||||
|
||||
# Restore the original device if you want to be a good citizen, though not strictly necessary
|
||||
# if the execution context is self-contained.
|
||||
if original_module_device is not None:
|
||||
setattr(loader_module, 'device', original_module_device)
|
||||
|
||||
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] WanVideo model loaded on {selected_device}")
|
||||
|
||||
# Return the model and the device for the sampler
|
||||
return (result[0], compute_device)
|
||||
else:
|
||||
logger.error("[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Could not find the module for WanVideoModelLoader to patch.")
|
||||
# Fallback to original behavior without patching
|
||||
return (original_loader.loadmodel(
|
||||
model, base_precision, load_device, quantization,
|
||||
compile_args, attention_mode, block_swap_args, lora,
|
||||
vram_management_args, extra_model=extra_model, vace_model=vace_model,
|
||||
fantasytalking_model=fantasytalking_model, multitalk_model=multitalk_model,
|
||||
fantasyportrait_model=fantasyportrait_model, rms_norm_function=rms_norm_function
|
||||
), compute_device)
|
||||
|
||||
logger.debug(
|
||||
f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Patching '{loader_module.__name__}'"
|
||||
)
|
||||
original_module_device = getattr(loader_module, "device", None)
|
||||
had_offload_attr = hasattr(loader_module, "offload_device")
|
||||
original_module_offload = getattr(loader_module, "offload_device", None)
|
||||
|
||||
setattr(loader_module, "device", selected_device)
|
||||
if swap_device_override is not None:
|
||||
setattr(loader_module, "offload_device", swap_device_override)
|
||||
elif compute_device == "cpu":
|
||||
setattr(loader_module, "offload_device", selected_device)
|
||||
|
||||
try:
|
||||
result = original_loader.loadmodel(
|
||||
model,
|
||||
base_precision,
|
||||
load_device,
|
||||
quantization,
|
||||
compile_args,
|
||||
attention_mode,
|
||||
normalized_block_swap if normalized_block_swap is not None else block_swap_args,
|
||||
lora,
|
||||
vram_management_args,
|
||||
extra_model=extra_model,
|
||||
vace_model=vace_model,
|
||||
fantasytalking_model=fantasytalking_model,
|
||||
multitalk_model=multitalk_model,
|
||||
fantasyportrait_model=fantasyportrait_model,
|
||||
rms_norm_function=rms_norm_function,
|
||||
)
|
||||
finally:
|
||||
if original_module_device is not None:
|
||||
setattr(loader_module, "device", original_module_device)
|
||||
if had_offload_attr:
|
||||
setattr(loader_module, "offload_device", original_module_offload)
|
||||
else:
|
||||
try:
|
||||
delattr(loader_module, "offload_device")
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
patcher = result[0]
|
||||
if normalized_block_swap is not None:
|
||||
try:
|
||||
transformer_options = patcher.model_options.setdefault("transformer_options", {})
|
||||
transformer_options["block_swap_args"] = normalized_block_swap
|
||||
except AttributeError:
|
||||
logger.warning(
|
||||
"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Unable to propagate normalized block swap settings"
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] WanVideo model loaded on %s with swap_device=%s",
|
||||
selected_device,
|
||||
str(swap_device_override) if swap_device_override is not None else "default",
|
||||
)
|
||||
|
||||
return (patcher, compute_device)
|
||||
|
||||
|
||||
class WanVideoSampler:
|
||||
@@ -613,39 +699,70 @@ class WanVideoSampler:
|
||||
|
||||
def process(self, model, compute_device, **kwargs):
|
||||
from . import set_current_device
|
||||
logger.info(
|
||||
f"[MultiGPU WanVideoSampler] Received request to process on: {compute_device}"
|
||||
)
|
||||
|
||||
logger.info(f"[MultiGPU WanVideoSampler] Received request to process on: {compute_device}")
|
||||
patcher = model
|
||||
transformer = None
|
||||
if hasattr(patcher, "model"):
|
||||
transformer = getattr(patcher.model, "diffusion_model", None)
|
||||
|
||||
# Resolve the target device and update the global sampler context
|
||||
target_device = None
|
||||
if compute_device:
|
||||
target_device = torch.device(compute_device)
|
||||
set_current_device(target_device)
|
||||
else:
|
||||
target_device = mm.get_torch_device()
|
||||
|
||||
normalized_swap_device = None
|
||||
transformer_options = {}
|
||||
if hasattr(patcher, "model_options"):
|
||||
transformer_options = patcher.model_options.get("transformer_options", {})
|
||||
block_swap_args = transformer_options.get("block_swap_args") if transformer_options else None
|
||||
if block_swap_args:
|
||||
swap_label = block_swap_args.get("resolved_swap_device") or block_swap_args.get("swap_device")
|
||||
if swap_label:
|
||||
try:
|
||||
normalized_swap_device = torch.device(str(swap_label))
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"[MultiGPU WanVideoSampler] Invalid swap device '%s', leaving sampler offload unchanged",
|
||||
swap_label,
|
||||
)
|
||||
normalized_swap_device = None
|
||||
|
||||
original_sampler = NODE_CLASS_MAPPINGS["WanVideoSampler"]()
|
||||
sampler_module = inspect.getmodule(original_sampler)
|
||||
|
||||
original_module_device = None
|
||||
original_module_offload = None
|
||||
had_sampler_offload_attr = False
|
||||
if sampler_module is not None:
|
||||
original_module_device = getattr(sampler_module, "device", None)
|
||||
had_sampler_offload_attr = hasattr(sampler_module, "offload_device")
|
||||
original_module_offload = getattr(sampler_module, "offload_device", None)
|
||||
setattr(sampler_module, "device", target_device)
|
||||
|
||||
# Align offload device when running on CPU so intermediate tensors stay colocated.
|
||||
if compute_device == "cpu":
|
||||
if normalized_swap_device is not None:
|
||||
setattr(sampler_module, "offload_device", normalized_swap_device)
|
||||
elif compute_device == "cpu":
|
||||
setattr(sampler_module, "offload_device", target_device)
|
||||
|
||||
if original_module_device != target_device:
|
||||
logger.debug(
|
||||
f"[MultiGPU WanVideoSampler] Patched sampler module device: {original_module_device} -> {target_device}"
|
||||
)
|
||||
else:
|
||||
logger.error("[MultiGPU WanVideoSampler] Unable to resolve sampler module for device patching.")
|
||||
|
||||
if transformer is not None and normalized_swap_device is not None:
|
||||
transformer.offload_device = normalized_swap_device
|
||||
transformer.cache_device = normalized_swap_device
|
||||
|
||||
try:
|
||||
# The original sampler will internally use mm.get_torch_device(), which is now correctly set.
|
||||
return original_sampler.process(model=model, **kwargs)
|
||||
return original_sampler.process(model=patcher, **kwargs)
|
||||
finally:
|
||||
if sampler_module is not None and original_module_device is not None:
|
||||
setattr(sampler_module, "device", original_module_device)
|
||||
if sampler_module is not None:
|
||||
if original_module_device is not None:
|
||||
setattr(sampler_module, "device", original_module_device)
|
||||
if had_sampler_offload_attr:
|
||||
setattr(sampler_module, "offload_device", original_module_offload)
|
||||
else:
|
||||
try:
|
||||
delattr(sampler_module, "offload_device")
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
Reference in New Issue
Block a user