diff --git a/nodes.py b/nodes.py index 8d23c61..b26ded8 100644 --- a/nodes.py +++ b/nodes.py @@ -186,33 +186,33 @@ class WanVideoTextEncode: cache_dir = os.path.join(script_directory, 'text_embed_cache') os.makedirs(cache_dir, exist_ok=True) - # Separate cache keys and paths for positive and negative - pos_cache_key = f"positive|{positive_prompt}" - pos_cache_hash = hashlib.sha256(pos_cache_key.encode('utf-8')).hexdigest() - pos_cache_path = os.path.join(cache_dir, f"{pos_cache_hash}.pt") - - neg_cache_key = f"negative|{negative_prompt}" - neg_cache_hash = hashlib.sha256(neg_cache_key.encode('utf-8')).hexdigest() - neg_cache_path = os.path.join(cache_dir, f"{neg_cache_hash}.pt") + # Use unified cache key for any prompt + def get_cache_path(prompt): + cache_key = prompt.strip() + cache_hash = hashlib.sha256(cache_key.encode('utf-8')).hexdigest() + return os.path.join(cache_dir, f"{cache_hash}.pt") context = None context_null = None + pos_cache_path = get_cache_path(positive_prompt) + neg_cache_path = get_cache_path(negative_prompt) + # Try to load positive prompt embeds if os.path.exists(pos_cache_path): try: - log.info(f"Loading positive prompt embeds from cache: {pos_cache_path}") + log.info(f"Loading prompt embeds from cache: {pos_cache_path}") context = torch.load(pos_cache_path) except Exception as e: - log.warning(f"Failed to load positive cache: {e}, will re-encode.") + log.warning(f"Failed to load cache: {e}, will re-encode.") # Try to load negative prompt embeds if os.path.exists(neg_cache_path): try: - log.info(f"Loading negative prompt embeds from cache: {neg_cache_path}") + log.info(f"Loading prompt embeds from cache: {neg_cache_path}") context_null = torch.load(neg_cache_path) except Exception as e: - log.warning(f"Failed to load negative cache: {e}, will re-encode.") + log.warning(f"Failed to load cache: {e}, will re-encode.") # If both loaded, return combined if context is not None and context_null is not None: @@ -277,15 +277,15 @@ class WanVideoTextEncode: try: if not os.path.exists(pos_cache_path): torch.save(context, pos_cache_path) - log.info(f"Saved positive prompt embeds to cache: {pos_cache_path}") + log.info(f"Saved prompt embeds to cache: {pos_cache_path}") except Exception as e: - log.warning(f"Failed to save positive cache: {e}") + log.warning(f"Failed to save cache: {e}") try: if not os.path.exists(neg_cache_path): torch.save(context_null, neg_cache_path) - log.info(f"Saved negative prompt embeds to cache: {neg_cache_path}") + log.info(f"Saved prompt embeds to cache: {neg_cache_path}") except Exception as e: - log.warning(f"Failed to save negative cache: {e}") + log.warning(f"Failed to save cache: {e}") return (prompt_embeds_dict,) @@ -319,7 +319,7 @@ class WanVideoTextEncodeSingle: "t5": ("WANTEXTENCODER",), "force_offload": ("BOOLEAN", {"default": True}), "model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}), - "cache_type": (["disabled", "positive", "negative", "both"], {"default": "positive", "tooltip": "Which cache to save to: positive, negative, or both"}), + "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"}), } } @@ -329,33 +329,24 @@ 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, cache_type="disabled"): - # Check for cache if enabled + def process(self, prompt, t5=None, force_offload=True, model_to_offload=None, use_disk_cache=False): + # Unified cache logic: use a single cache file per unique prompt encoded = None - if cache_type != "disabled": + if use_disk_cache: cache_dir = os.path.join(script_directory, 'text_embed_cache') os.makedirs(cache_dir, exist_ok=True) - pos_cache_key = f"positive|{prompt}" - pos_cache_hash = hashlib.sha256(pos_cache_key.encode('utf-8')).hexdigest() - pos_cache_path = os.path.join(cache_dir, f"{pos_cache_hash}.pt") - neg_cache_key = f"neg|{prompt}" - neg_cache_hash = hashlib.sha256(neg_cache_key.encode('utf-8')).hexdigest() - neg_cache_path = os.path.join(cache_dir, f"{neg_cache_hash}.pt") - # Try positive cache first - if os.path.exists(pos_cache_path): + def get_cache_path(prompt): + cache_key = prompt.strip() + cache_hash = hashlib.sha256(cache_key.encode('utf-8')).hexdigest() + return os.path.join(cache_dir, f"{cache_hash}.pt") + cache_path = get_cache_path(prompt) + if os.path.exists(cache_path): try: - log.info(f"Loading prompt embeds from positive cache: {pos_cache_path}") - encoded = torch.load(pos_cache_path) + log.info(f"Loading prompt embeds from cache: {cache_path}") + encoded = torch.load(cache_path) except Exception as e: - log.warning(f"Failed to load positive cache: {e}, will re-encode.") - # Try negative cache if not found in positive - if encoded is None and os.path.exists(neg_cache_path): - try: - log.info(f"Loading prompt embeds from negative cache: {neg_cache_path}") - encoded = torch.load(neg_cache_path) - except Exception as e: - log.warning(f"Failed to load negative cache: {e}, will re-encode.") - + log.warning(f"Failed to load cache: {e}, will re-encode.") + if t5 is None and encoded is None: raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.") @@ -377,22 +368,14 @@ class WanVideoTextEncodeSingle: encoder.model.to(offload_device) mm.soft_empty_cache() - # Save to selected cache(s) if enabled + # Save to cache if enabled if use_disk_cache: - if cache_type in ("positive", "both"): - try: - if not os.path.exists(pos_cache_path): - torch.save(encoded, pos_cache_path) - log.info(f"Saved prompt embeds to positive cache: {pos_cache_path}") - except Exception as e: - log.warning(f"Failed to save positive cache: {e}") - if cache_type in ("negative", "both"): - try: - if not os.path.exists(neg_cache_path): - torch.save(encoded, neg_cache_path) - log.info(f"Saved prompt embeds to negative cache: {neg_cache_path}") - except Exception as e: - log.warning(f"Failed to save negative cache: {e}") + try: + if not os.path.exists(cache_path): + torch.save(encoded, cache_path) + log.info(f"Saved prompt embeds to cache: {cache_path}") + except Exception as e: + log.warning(f"Failed to save cache: {e}") prompt_embeds_dict = { "prompt_embeds": encoded,