Add optional prompt disk caching

Caches text embeds to disk with unique hash based on T5 name, dtype and prompts. If found on disk skips encoding, persists after restarts.
This commit is contained in:
kijai
2025-07-23 00:52:56 +03:00
parent 729b6fdf7c
commit fa93f5b3c1
2 changed files with 39 additions and 8 deletions
+38 -8
View File
@@ -4,7 +4,7 @@ import torch.nn.functional as F
import numpy as np
from tqdm import tqdm
import inspect
import hashlib
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from .wanvideo.modules.model import rope_params
@@ -166,6 +166,7 @@ class WanVideoTextEncode:
"optional": {
"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"}),
}
}
@@ -175,7 +176,29 @@ class WanVideoTextEncode:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length"
def process(self, t5, positive_prompt, negative_prompt,force_offload=True, model_to_offload=None):
def process(self, t5, positive_prompt, negative_prompt, force_offload=True, model_to_offload=None, use_disk_cache=False):
if use_disk_cache:
# Prepare cache directory
cache_dir = os.path.join(script_directory, 'text_embed_cache')
os.makedirs(cache_dir, exist_ok=True)
# Build a unique cache key from all relevant inputs
encoder_id = str(t5["name"])
dtype_str = str(t5["dtype"])
cache_key = f"{encoder_id}|{positive_prompt}|{negative_prompt}|{dtype_str}"
cache_hash = hashlib.sha256(cache_key.encode('utf-8')).hexdigest()
cache_path = os.path.join(cache_dir, f"{cache_hash}.pt")
# Try to load from cache
if os.path.exists(cache_path):
try:
log.info(f"Loading text embeds from cache: {cache_path}")
prompt_embeds_dict = torch.load(cache_path)
return (prompt_embeds_dict,)
except Exception as e:
log.warning(f"Failed to load cache: {e}, will re-encode.")
if model_to_offload is not None:
log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device)
@@ -188,14 +211,13 @@ class WanVideoTextEncode:
positive_prompts_raw = [p.strip() for p in positive_prompt.split('|')]
positive_prompts = []
all_weights = []
for p in positive_prompts_raw:
cleaned_prompt, weights = self.parse_prompt_weights(p)
positive_prompts.append(cleaned_prompt)
all_weights.append(weights)
encoder.model.to(device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=True):
context = encoder(positive_prompts, device)
context_null = encoder([negative_prompt], device)
@@ -212,9 +234,17 @@ class WanVideoTextEncode:
mm.soft_empty_cache()
prompt_embeds_dict = {
"prompt_embeds": context,
"negative_prompt_embeds": context_null,
}
"prompt_embeds": context,
"negative_prompt_embeds": context_null,
}
if use_disk_cache:
try:
torch.save(prompt_embeds_dict, cache_path)
log.info(f"Saved text embeds to cache: {cache_path}")
except Exception as e:
log.warning(f"Failed to save cache: {e}")
return (prompt_embeds_dict,)
def parse_prompt_weights(self, prompt):
+1
View File
@@ -1398,6 +1398,7 @@ class LoadWanVideoT5TextEncoder:
text_encoder = {
"model": T5_text_encoder,
"dtype": dtype,
"name": model_name,
}
return (text_encoder,)