From ee6e790d80c7e61414b9ce44dc777da8f52d51d2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 28 Jul 2025 09:50:50 +0300 Subject: [PATCH] Update nodes.py --- nodes.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index 3b38c6b..7864a9e 100644 --- a/nodes.py +++ b/nodes.py @@ -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):