When dropping out, caching, and doing DOP, make sure we select a cahced trigger word when dropped out so DOP matches the drop out embeddings.

This commit is contained in:
Jaret Burkett
2026-08-11 06:59:20 -06:00
parent 8d4beedd04
commit f421542df4
2 changed files with 55 additions and 2 deletions

View File

@@ -1,5 +1,6 @@
import os
import torch
from safetensors import safe_open
from safetensors.torch import load_file, save_file
@@ -137,6 +138,8 @@ class AdvancedPromptEmbeds:
def save(self, path):
data = {}
metadata = {"class_name": self.__class__.__name__}
if self._frozen_dtype_keys:
metadata["frozen_dtype_keys"] = ",".join(self._frozen_dtype_keys)
for key, value in self._store.items():
if len(value) != 1:
raise ValueError(
@@ -150,6 +153,8 @@ class AdvancedPromptEmbeds:
def load(cls, path=None):
if path is not None:
loaded = load_file(path)
with safe_open(path, framework="pt") as f:
metadata = f.metadata()
else:
raise ValueError("Must provide a path")
@@ -157,7 +162,18 @@ class AdvancedPromptEmbeds:
for key in loaded.keys():
data[key] = loaded[key]
return cls(**data)
pe = cls(**data)
frozen = []
if metadata is not None and metadata.get("frozen_dtype_keys"):
frozen = metadata["frozen_dtype_keys"].split(",")
# files cached before frozen keys were saved have no metadata for them.
# Integer/bool tensors (token tags, etc) must never be dtype cast, so
# freeze them regardless.
for key, value in pe._store.items():
if key not in frozen and not value[0].is_floating_point():
frozen.append(key)
pe.frozen_dtype_keys = frozen
return pe
@classmethod
def concat_prompt_embeds(

View File

@@ -2092,7 +2092,10 @@ class TextEmbeddingFileItemDTOMixin:
self._dop_text_embedding_path: Union[str, None] = None
# blank caption embeds used for caption dropout when caching text embeddings
self._blank_text_embedding_path: Union[str, None] = None
# DOP embeds for dropout steps (dropout caption with trigger replaced by class)
self._dop_blank_text_embedding_path: Union[str, None] = None
self._loaded_text_embedding_path: Union[str, None] = None
self._caption_was_dropped = False
self.is_text_embedding_cached = False
self.text_embedding_load_device = 'cpu'
self.text_embedding_version = 1
@@ -2161,6 +2164,28 @@ class TextEmbeddingFileItemDTOMixin:
return inject_trigger_into_prompt('', trigger=self.trigger_word, add_if_not_present=True)
return ''
def get_dop_dropout_caption(self: 'FileItemDTO'):
# live encoding replaces the trigger word with the preservation class on the
# dropped caption (class only), so the cached DOP dropout caption must match
dropout_caption = self.get_dropout_caption()
if self.trigger_word is not None:
return dropout_caption.replace(
self.trigger_word, self.dataset_config.diff_output_preservation_class
)
return dropout_caption
def get_dop_blank_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._dop_blank_text_embedding_path is not None and not recalculate:
return self._dop_blank_text_embedding_path
else:
# if the DOP dropout caption matches the dropout caption, this hashes to
# the same path as the blank embedding and the cache file is shared
self._dop_blank_text_embedding_path = self._build_text_embedding_path(
caption_override=self.get_dop_dropout_caption()
)
return self._dop_blank_text_embedding_path
def get_blank_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._blank_text_embedding_path is not None and not recalculate:
return self._blank_text_embedding_path
@@ -2185,17 +2210,24 @@ class TextEmbeddingFileItemDTOMixin:
return
if self.prompt_embeds is None:
te_path = self.get_text_embedding_path()
self._caption_was_dropped = False
if self.dataset_config.caption_dropout_rate > 0:
# get a random float form 0 to 1
rand = random.random()
if rand < self.dataset_config.caption_dropout_rate:
# drop the caption by using the cached blank embedding
te_path = self.get_blank_text_embedding_path()
self._caption_was_dropped = True
# load it from disk
self.prompt_embeds = PromptEmbeds.load(te_path)
self._loaded_text_embedding_path = te_path
if self.dataset_config.diff_output_preservation and self.dop_prompt_embeds is None:
dop_path = self.get_dop_text_embedding_path()
if self._caption_was_dropped:
# match live encoding, which builds the DOP caption from the
# dropped caption (trigger word replaced with the class)
dop_path = self.get_dop_blank_text_embedding_path()
else:
dop_path = self.get_dop_text_embedding_path()
if dop_path == self._loaded_text_embedding_path:
# no trigger word in caption, same embedding
self.dop_prompt_embeds = self.prompt_embeds
@@ -2234,6 +2266,11 @@ class TextEmbeddingCachingMixin:
if blank_path != text_embedding_path:
# cache the dropout caption embedding (blank, or trigger word only)
encode_targets.append((blank_path, file_item.get_dropout_caption()))
if self.dataset_config.diff_output_preservation:
# cache the DOP version of the dropout caption (class only)
dop_blank_path = file_item.get_dop_blank_text_embedding_path(recalculate=True)
if dop_blank_path not in [t[0] for t in encode_targets] + [text_embedding_path]:
encode_targets.append((dop_blank_path, file_item.get_dop_dropout_caption()))
# only process if not saved to disk
encode_targets = [t for t in encode_targets if not os.path.exists(t[0])]
if len(encode_targets) > 0: