feat: add WanVideoTextEncodeCached and WanVideoTextEncodeSingle classes for enhanced text encoding functionality

This commit is contained in:
John Pollock
2025-10-08 06:08:16 -05:00
parent d300c12e9f
commit 04013c3ee5
2 changed files with 97 additions and 29 deletions
+4
View File
@@ -213,6 +213,8 @@ from .nodes import (
from .wanvideo import (
LoadWanVideoT5TextEncoder,
WanVideoTextEncode,
WanVideoTextEncodeCached,
WanVideoTextEncodeSingle,
WanVideoVAELoader,
WanVideoTinyVAELoader,
WanVideoBlockSwap,
@@ -358,6 +360,8 @@ 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,
+93 -29
View File
@@ -1,12 +1,3 @@
"""WanVideoWrapper integration helpers.
For the current progress checklist and outstanding tasks, see
`.github/instructions/ComfyUI-MultiGPU.instructions.md`.
"""
import logging
import torch
import sys
@@ -24,25 +15,6 @@ from accelerate import init_empty_weights
import os
import importlib.util
scheduler_list = [
"unipc", "unipc/beta",
"dpm++", "dpm++/beta",
"dpm++_sde", "dpm++_sde/beta",
"euler", "euler/beta",
"deis",
"lcm", "lcm/beta",
"res_multistep",
"flowmatch_causvid",
"flowmatch_distill",
"flowmatch_pusa",
"multitalk",
"sa_ode_stable"
]
rope_functions = ["default", "comfy", "comfy_chunked"]
logger = logging.getLogger("MultiGPU")
class LoadWanVideoT5TextEncoder:
@@ -90,6 +62,55 @@ class LoadWanVideoT5TextEncoder:
return text_encoder, device
class WanVideoTextEncodeCached:
@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("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/text_encoders'"}),
"precision": (["fp32", "bf16"], {"default": "bf16"}),
"positive_prompt": ("STRING", {"default": "", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
"quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"use_disk_cache": ("BOOLEAN", {"default": True, "tooltip": "Cache the text embeddings to disk for faster re-use, under the custom_nodes/ComfyUI-WanVideoWrapper/text_embed_cache directory"}),
"load_device": (devices, {"default": default_device}
),
},
"optional": {
"extender_args": ("WANVIDEOPROMPTEXTENDER_ARGS", {"tooltip": "Use this node to extend the prompt with additional text."}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", "WANVIDEOTEXTEMBEDS", "STRING")
RETURN_NAMES = ("text_embeds", "negative_text_embeds", "positive_prompt")
OUTPUT_TOOLTIPS = ("The text embeddings for both prompts", "The text embeddings for the negative prompt only (for NAG)", "Positive prompt to display prompt extender results")
FUNCTION = "process"
CATEGORY = "multigpu/WanVideoWrapper"
DESCRIPTION = """Encodes text prompts into text embeddings. This node loads and completely unloads the T5 after done, leaving no VRAM or RAM imprint."""
def process(self, model_name, precision, positive_prompt, negative_prompt, quantization='disabled', use_disk_cache=True, load_device=None, extender_args=None):
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"
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeCachedMulitiGPU] current_device set to: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeCachedMulitiGPU] device set to: {device}")
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncodeCached"]()
prompt_embeds_dict, negative_text_embeds, positive_prompt_out = original_encoder.process(model_name, precision, positive_prompt, negative_prompt, quantization, use_disk_cache, device, extender_args)
return prompt_embeds_dict, negative_text_embeds, positive_prompt_out
class WanVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
@@ -103,7 +124,6 @@ class WanVideoTextEncode:
"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"}),
#"device": (["gpu", "cpu"], {"default": "gpu", "tooltip": "Device to run the text encoding on."}),
}
}
@@ -141,6 +161,50 @@ class WanVideoTextEncode:
original_parser = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
return original_parser.parse_prompt_weights(prompt)
class WanVideoTextEncodeSingle:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"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 prompt into text embedding."
def process(self, 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][WanVideoTextEncodeSingleMulitiGPU] current_device set to: {load_device}")
logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeSingleMulitiGPU] device set to: {device}")
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncodeSingle"]()
prompt_embeds_dict = original_encoder.process(prompt, text_encoder, force_offload, model_to_offload, use_disk_cache, device)
return (prompt_embeds_dict)
class WanVideoVAELoader:
@classmethod
def INPUT_TYPES(s):