Update nodes.py
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user