Update nodes.py

This commit is contained in:
kijai
2025-07-28 09:50:50 +03:00
parent 5e725ff8fb
commit ee6e790d80
+10 -3
View File
@@ -229,7 +229,7 @@ class WanVideoTextEncode:
if t5 is None:
raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.")
if model_to_offload is not None:
if model_to_offload is not None and device == "gpu":
log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device)
@@ -344,6 +344,7 @@ class WanVideoTextEncodeSingle:
"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."}),
}
}
@@ -353,7 +354,7 @@ class WanVideoTextEncodeSingle:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes text prompt into text embedding."
def process(self, prompt, t5=None, force_offload=True, model_to_offload=None, use_disk_cache=False):
def process(self, prompt, t5=None, force_offload=True, model_to_offload=None, use_disk_cache=False, device="gpu"):
# Unified cache logic: use a single cache file per unique prompt
encoded = None
if use_disk_cache:
@@ -375,7 +376,7 @@ class WanVideoTextEncodeSingle:
raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.")
if encoded is None:
if model_to_offload is not None:
if model_to_offload is not None and device == "gpu":
log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device)
mm.soft_empty_cache()
@@ -383,6 +384,12 @@ class WanVideoTextEncodeSingle:
encoder = t5["model"]
dtype = t5["dtype"]
if device == "gpu":
device = mm.get_torch_device()
encoder.model.to(device)
elif device == "cpu":
encoder.model.to(torch.device("cpu"))
encoder.model.to(device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=True):