Add LoadWanVideoClipTextEncoder and WanVideoClipVisionEncode for multi-GPU support

This commit is contained in:
John Pollock
2025-10-09 19:05:28 -05:00
parent 2f605e0db3
commit 375cbd093c
2 changed files with 84 additions and 0 deletions
+4
View File
@@ -224,6 +224,8 @@ from .wanvideo import (
WanVideoSampler,
WanVideoVACEEncode,
WanVideoEncode,
LoadWanVideoClipTextEncoder,
WanVideoClipVisionEncode,
)
from .wrappers import (
@@ -373,6 +375,8 @@ wanvideo_nodes = {
"WanVideoSamplerMultiGPU": WanVideoSampler,
"WanVideoVACEEncodeMultiGPU": WanVideoVACEEncode,
"WanVideoEncodeMultiGPU": WanVideoEncode,
"LoadWanVideoClipTextEncoderMultiGPU": LoadWanVideoClipTextEncoder,
"WanVideoClipVisionEncodeMultiGPU": WanVideoClipVisionEncode,
}
register_and_count(["ComfyUI-WanVideoWrapper", "comfyui-wanvideowrapper"], wanvideo_nodes)
+80
View File
@@ -261,6 +261,44 @@ class LoadWanVideoT5TextEncoder:
return text_encoder, device
class LoadWanVideoClipTextEncoder:
@classmethod
def INPUT_TYPES(s):
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
return {
"required": {
"model_name": (folder_paths.get_filename_list("clip_vision") + folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/clip_vision'"}),
"precision": (["fp16", "fp32", "bf16"],
{"default": "fp16"}
),
},
"optional": {
"device": (devices, {"default": default_device}),
}
}
RETURN_TYPES = ("CLIP_VISION", "MULTIGPUDEVICE")
RETURN_NAMES = ("wan_clip_vision", "load_device")
FUNCTION = "loadmodel"
CATEGORY = "multigpu/WanVideoWrapper"
DESCRIPTION = "Loads Wan clip_vision model from 'ComfyUI/models/clip_vision'"
def loadmodel(self, model_name, precision, device=None):
from . import set_current_device
set_current_device(device)
if device == "cpu":
load_device = "offload_device"
else:
load_device = "main_device"
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoder"]()
clip_model = original_loader.loadmodel(model_name, precision, load_device)
return clip_model, device
class WanVideoTextEncodeCached:
@classmethod
def INPUT_TYPES(s):
@@ -632,3 +670,45 @@ class WanVideoEncode:
return original_encode.encode(vae[0], image, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, noise_aug_strength, latent_strength, mask)
finally:
encode_module.device = original_module_device
class WanVideoClipVisionEncode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip_vision": ("CLIP_VISION",),
"load_device": ("MULTIGPUDEVICE",),
"image_1": ("IMAGE", {"tooltip": "Image to encode"}),
"strength_1": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional clip embed multiplier"}),
"strength_2": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional clip embed multiplier"}),
"crop": (["center", "disabled"], {"default": "center", "tooltip": "Crop image to 224x224 before encoding"}),
"combine_embeds": (["average", "sum", "concat", "batch"], {"default": "average", "tooltip": "Method to combine multiple clip embeds"}),
"force_offload": ("BOOLEAN", {"default": True}),
},
"optional": {
"image_2": ("IMAGE", ),
"negative_image": ("IMAGE", {"tooltip": "image to use for uncond"}),
"tiles": ("INT", {"default": 0, "min": 0, "max": 16, "step": 2, "tooltip": "Use matteo's tiled image encoding for improved accuracy"}),
"ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Ratio of the tile average"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_CLIPEMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "multigpu/WanVideoWrapper"
def process(self, clip_vision, load_device, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None, negative_image=None, tiles=0, ratio=1.0):
from . import set_current_device
original_encode = NODE_CLASS_MAPPINGS["WanVideoClipVisionEncode"]()
encode_module = inspect.getmodule(original_encode)
original_module_device = encode_module.device
set_current_device(load_device)
compute_device_to_be_patched = mm.get_torch_device()
encode_module.device = compute_device_to_be_patched
try:
return original_encode.process(clip_vision[0], image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2, negative_image, tiles, ratio)
finally:
encode_module.device = original_module_device