This commit is contained in:
John Pollock
2025-10-08 23:06:54 -05:00
parent 04013c3ee5
commit 7c87983d72
+129 -254
View File
@@ -17,6 +17,123 @@ import importlib.util
logger = logging.getLogger("MultiGPU")
class WanVideoModelLoader:
@classmethod
def INPUT_TYPES(s):
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
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_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"default": "disabled",
"tooltip": "Optional quantization method, 'disabled' acts as autoselect based by weights. Scaled modes only work with matching weights, _fast modes (fp8 matmul) require CUDA compute capability >= 8.9 (NVIDIA 4000 series and up), e4m3fn generally can not be torch.compiled on compute capability < 8.9 (3000 series and under)"}),
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
"compute_device": (devices, {"default": default_device}),
},
"optional": {
"attention_mode": ([
"sdpa",
"flash_attn_2",
"flash_attn_3",
"sageattn",
"sageattn_3",
"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"}),
"extra_model": ("VACEPATH", {"default": None, "tooltip": "Extra model to add to the main model, ie. VACE or MTV Crafter"}),
"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"}),
"rms_norm_function": (["default", "pytorch"], {"default": "default", "tooltip": "RMSNorm function to use, 'pytorch' is the new native torch RMSNorm, which is faster (when not using torch.compile mostly) but changes results slightly. 'default' is the original WanRMSNorm"}),
}
}
RETURN_TYPES = ("WANVIDEOMODEL", "MULTIGPUDEVICE",)
RETURN_NAMES = ("model", "compute_device",)
FUNCTION = "loadmodel"
CATEGORY = "multigpu/WanVideoWrapper"
def loadmodel(self, model, base_precision, compute_device, quantization, load_device,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None,
vram_management_args=None, extra_model=None, vace_model=None,
fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None,
rms_norm_function="default"):
from . import set_current_device
set_current_device(compute_device)
compute_device_to_be_patched = mm.get_torch_device()
original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]()
loader_module = inspect.getmodule(original_loader)
original_module_device = loader_module.device
loader_module.device = compute_device_to_be_patched
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,)
loader_module.device = original_module_device
patcher = result[0]
return (patcher, compute_device)
class WanVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive_prompt": ("STRING", {"default": "", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
},
"optional": {
"t5": ("WANTEXTENCODER",),
"load_device": ("MULTIGPUDEVICE",),
"force_offload": ("BOOLEAN", {"default": True}),
"model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}),
"use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use, under the custom_nodes/ComfyUI-WanVideoWrapper/text_embed_cache directory"}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "multigpu/WanVideoWrapper"
DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length"
def process(self, positive_prompt, negative_prompt, t5=None, load_device=None,force_offload=True, model_to_offload=None, use_disk_cache=False):
from . import set_current_device
set_current_device(load_device)
if load_device == "cpu":
device = "cpu"
else:
device = "gpu"
if t5 is not None:
text_encoder = t5[0]
else:
text_encoder = None
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] current_device set to: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] device set to: {device}")
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
prompt_embeds_dict = original_encoder.process(positive_prompt, negative_prompt, text_encoder, force_offload, model_to_offload, use_disk_cache, device)
return (prompt_embeds_dict)
def parse_prompt_weights(self, prompt):
"""Extract text and weights from prompts with (text:weight) format"""
original_parser = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
return original_parser.parse_prompt_weights(prompt)
class LoadWanVideoT5TextEncoder:
@classmethod
def INPUT_TYPES(s):
@@ -58,7 +175,6 @@ class LoadWanVideoT5TextEncoder:
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoder"]()
text_encoder = original_loader.loadmodel(model_name, precision, load_device, quantization)
# Return both the text encoder AND the selected device
return text_encoder, device
@@ -110,57 +226,6 @@ class WanVideoTextEncodeCached:
return prompt_embeds_dict, negative_text_embeds, positive_prompt_out
class WanVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive_prompt": ("STRING", {"default": "", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
},
"optional": {
"t5": ("WANTEXTENCODER",),
"load_device": ("MULTIGPUDEVICE",),
"force_offload": ("BOOLEAN", {"default": True}),
"model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}),
"use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use, under the custom_nodes/ComfyUI-WanVideoWrapper/text_embed_cache directory"}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "multigpu/WanVideoWrapper"
DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length"
def process(self, positive_prompt, negative_prompt, t5=None, load_device=None,force_offload=True, model_to_offload=None, use_disk_cache=False):
from . import set_current_device
if load_device is not None:
set_current_device(load_device)
if load_device == "cpu":
device = "cpu"
else:
device = "gpu"
if t5 is not None:
text_encoder = t5[0]
else:
text_encoder = None
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] current_device set to: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] device set to: {device}")
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
prompt_embeds_dict = original_encoder.process(positive_prompt, negative_prompt, text_encoder, force_offload, model_to_offload, use_disk_cache, device)
return (prompt_embeds_dict)
def parse_prompt_weights(self, prompt):
"""Extract text and weights from prompts with (text:weight) format"""
original_parser = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
return original_parser.parse_prompt_weights(prompt)
class WanVideoTextEncodeSingle:
@classmethod
def INPUT_TYPES(s):
@@ -235,7 +300,7 @@ class WanVideoVAELoader:
if load_device is not None:
set_current_device(load_device)
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoVAELoader] load_device set to: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoVAELoaderMultiGPU] load_device set to: {load_device}")
original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]()
vae_model = original_loader.loadmodel(model_name, precision, compile_args)
@@ -345,8 +410,7 @@ class WanVideoImageToVideoEncode:
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, load_device=None):
from . import set_current_device
if load_device is not None:
set_current_device(load_device)
set_current_device(load_device)
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] load device: {load_device}")
@@ -537,214 +601,25 @@ class WanVideoDecode:
def decode(self, vae, load_device, samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization="default"):
from . import set_current_device
if load_device is not None:
set_current_device(load_device)
set_current_device(load_device)
compute_device_to_be_patched = mm.get_torch_device()
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] load device: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoDecodeMultiGPU] load device: {load_device}")
device = mm.get_torch_device()
PATCH_SIZE = (1, 2, 2)
offload_device = mm.unet_offload_device()
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] torch device: {device}")
if vae is not None:
vae = vae[0]
mm.soft_empty_cache()
video = samples.get("video", None)
if video is not None:
video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0)
return video.cpu().float(),
latents = samples["samples"]
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
drop_last = samples.get("drop_last", False)
is_looped = samples.get("looped", False)
vae.to(device)
latents = latents.to(device = device, dtype = vae.dtype)
mm.soft_empty_cache()
if has_ref:
latents = latents[:, :, 1:]
if drop_last:
latents = latents[:, :, :-1]
if type(vae).__name__ == "TAEHV":
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4))[0].permute(1, 0, 2, 3)
images = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).cpu().float()
return (images,)
else:
if end_image is not None:
enable_vae_tiling = False
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
images = images.cpu().float()
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped:
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
temp_images = temp_images.cpu().float()
temp_images = (temp_images - temp_images.min()) / (temp_images.max() - temp_images.min())
images = torch.cat([temp_images[:, 9:].to(images), images[:, 5:]], dim=1)
if end_image is not None:
images = images[:, 0:-1]
vae.to(offload_device)
mm.soft_empty_cache()
images.clamp_(0.0, 1.0)
return (images.permute(1, 2, 3, 0),)
class WanVideoModelLoader:
@classmethod
def INPUT_TYPES(s):
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
# Get the original node's input types to stay up-to-date
original_types = NODE_CLASS_MAPPINGS["WanVideoModelLoader"].INPUT_TYPES()
# Update with our custom device selection
original_types["required"]["compute_device"] = (devices, {"default": default_device})
return original_types
RETURN_TYPES = ("WANVIDEOMODEL", "MULTIGPUDEVICE",)
RETURN_NAMES = ("model", "compute_device",)
FUNCTION = "loadmodel"
CATEGORY = "multigpu/WanVideoWrapper"
def loadmodel(self, model, base_precision, compute_device, quantization, load_device,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None,
vram_management_args=None, extra_model=None, vace_model=None,
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}"
)
selected_device = torch.device(compute_device)
set_current_device(selected_device)
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"]()
original_loader = NODE_CLASS_MAPPINGS["WanVideoDecode"]()
loader_module = inspect.getmodule(original_loader)
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,
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,
)
return (result[0], compute_device)
original_module_device = loader_module.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)
loader_module.device = compute_device_to_be_patched
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)
result = original_loader.decode(vae[0], samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization)
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
loader_module.device = original_module_device
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"
)
decode = result[0]
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)
return (decode,)
class WanVideoSampler: