Rework DOP so it works with caching text embeddings

This commit is contained in:
Jaret Burkett
2026-08-09 22:13:49 -06:00
parent 5ff8a0435a
commit 257da9b586
5 changed files with 138 additions and 30 deletions

View File

@@ -924,6 +924,10 @@ class DatasetConfig:
self.default_caption: str = kwargs.get('default_caption', None)
# trigger word for just this dataset
self.trigger_word: str = kwargs.get('trigger_word', None)
# set automatically from the train config when diff output preservation is enabled.
# the dataset trigger word is replaced with the class in the caption for DOP embeddings
self.diff_output_preservation: bool = kwargs.get('diff_output_preservation', False)
self.diff_output_preservation_class: str = kwargs.get('diff_output_preservation_class', '')
random_triggers = kwargs.get('random_triggers', [])
# if they are a string, load them from a file
if isinstance(random_triggers, str) and os.path.exists(random_triggers):
@@ -1489,11 +1493,6 @@ def validate_configs(
# see if any datasets are caching text embeddings
is_caching_text_embeddings = any(dataset.cache_text_embeddings for dataset in dataset_configs)
if is_caching_text_embeddings:
# check if they are doing differential output preservation
if train_config.diff_output_preservation:
raise ValueError("Cannot use differential output preservation with caching text embeddings. Please set diff_output_preservation to False.")
# make sure they are all cached
for dataset in dataset_configs:
if not dataset.cache_text_embeddings:

View File

@@ -315,6 +315,8 @@ class DataLoaderBatchDTO:
)
self.prompt_embeds: Union[PromptEmbeds, None] = None
# diff output preservation embeds (trigger word replaced with class)
self.dop_prompt_embeds: Union[PromptEmbeds, None] = None
# if self.file_items[0].control_tensor is not None:
# if any have a control tensor, we concatenate them
if any([x.control_tensor is not None for x in self.file_items]):
@@ -482,9 +484,31 @@ class DataLoaderBatchDTO:
y.text_embeds = [y.text_embeds]
prompt_embeds_list.append(y)
padding_side = self.file_items[0].te_padding_side
self.prompt_embeds = concat_prompt_embeds(prompt_embeds_list, padding_side=padding_side)
if any([x.dop_prompt_embeds is not None for x in self.file_items]):
# find one to use as a base
base_dop_prompt_embeds = None
for x in self.file_items:
if x.dop_prompt_embeds is not None:
base_dop_prompt_embeds = x.dop_prompt_embeds
break
dop_prompt_embeds_list = []
for x in self.file_items:
if x.dop_prompt_embeds is None:
y = base_dop_prompt_embeds
else:
y = x.dop_prompt_embeds
if x.text_embedding_space_version == "zimage":
# z image needs to be a list if it is not already
if not isinstance(y.text_embeds, list):
y.text_embeds = [y.text_embeds]
dop_prompt_embeds_list.append(y)
padding_side = self.file_items[0].te_padding_side
self.dop_prompt_embeds = concat_prompt_embeds(dop_prompt_embeds_list, padding_side=padding_side)
if any([x.audio_tensor is not None for x in self.file_items]):
# find one to use as a base
base_audio_tensor = None

View File

@@ -321,6 +321,8 @@ class CaptionProcessingDTOMixin:
self.raw_caption_short: str = None
self.caption: str = None
self.caption_short: str = None
# caption with the trigger word replaced by the diff output preservation class
self.caption_dop: str = None
dataset_config: DatasetConfig = kwargs.get('dataset_config', None)
self.extra_values: List[float] = dataset_config.extra_values
@@ -367,6 +369,14 @@ class CaptionProcessingDTOMixin:
self.caption = self.get_caption()
if self.raw_caption_short is not None:
self.caption_short = self.get_caption(short_caption=True)
if self.dataset_config.diff_output_preservation:
# replace this dataset's trigger word with the preservation class.
# do it on the final caption so token order matches the normal caption
self.caption_dop = self.caption
if self.trigger_word is not None:
self.caption_dop = self.caption.replace(
self.trigger_word, self.dataset_config.diff_output_preservation_class
)
def get_caption(
self: 'FileItemDTO',
@@ -2077,17 +2087,20 @@ class TextEmbeddingFileItemDTOMixin:
super().__init__(*args, **kwargs)
self.prompt_embeds: Union[PromptEmbeds, None] = None
self._text_embedding_path: Union[str, None] = None
# diff output preservation embeds (caption with trigger word replaced by class)
self.dop_prompt_embeds: Union[PromptEmbeds, None] = None
self._dop_text_embedding_path: Union[str, None] = None
self.is_text_embedding_cached = False
self.text_embedding_load_device = 'cpu'
self.text_embedding_version = 1
def get_text_embedding_info_dict(self: 'FileItemDTO'):
def get_text_embedding_info_dict(self: 'FileItemDTO', caption_override=None):
# make sure the caption is loaded here
# TODO: we need a way to cache all the other features like trigger words, DOP, etc. For now, we need to throw an error if not compatible.
if self.caption is None:
self.load_caption()
item = OrderedDict([
("caption", self.caption),
("caption", self.caption if caption_override is None else caption_override),
("text_embedding_space_version", self.text_embedding_space_version),
("text_embedding_version", self.text_embedding_version),
])
@@ -2103,27 +2116,47 @@ class TextEmbeddingFileItemDTOMixin:
item["first_frame_in_te"] = True
return item
def _build_text_embedding_path(self: 'FileItemDTO', caption_override=None):
# we store text embeddings in a folder in same path as image called _text_embedding_cache
img_dir = os.path.dirname(self.path)
te_dir = os.path.join(img_dir, '_t_e_cache')
hash_dict = self.get_text_embedding_info_dict(caption_override=caption_override)
filename_no_ext = os.path.splitext(os.path.basename(self.path))[0]
# get base64 hash of md5 checksum of hash_dict
hash_input = json.dumps(hash_dict, sort_keys=True).encode('utf-8')
hash_str = base64.urlsafe_b64encode(hashlib.md5(hash_input).digest()).decode('ascii')
hash_str = hash_str.replace('=', '')
return os.path.join(te_dir, f'{filename_no_ext}_{hash_str}.safetensors')
def get_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._text_embedding_path is not None and not recalculate:
return self._text_embedding_path
else:
# we store text embeddings in a folder in same path as image called _text_embedding_cache
img_dir = os.path.dirname(self.path)
te_dir = os.path.join(img_dir, '_t_e_cache')
hash_dict = self.get_text_embedding_info_dict()
filename_no_ext = os.path.splitext(os.path.basename(self.path))[0]
# get base64 hash of md5 checksum of hash_dict
hash_input = json.dumps(hash_dict, sort_keys=True).encode('utf-8')
hash_str = base64.urlsafe_b64encode(hashlib.md5(hash_input).digest()).decode('ascii')
hash_str = hash_str.replace('=', '')
self._text_embedding_path = os.path.join(te_dir, f'{filename_no_ext}_{hash_str}.safetensors')
self._text_embedding_path = self._build_text_embedding_path()
return self._text_embedding_path
def get_dop_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._dop_text_embedding_path is not None and not recalculate:
return self._dop_text_embedding_path
else:
# make sure the caption is loaded so caption_dop is built
if self.caption is None:
self.load_caption()
# if the trigger word is not in the caption, this hashes to the same
# path as the normal embedding and the cache file is shared
self._dop_text_embedding_path = self._build_text_embedding_path(
caption_override=self.caption_dop
)
return self._dop_text_embedding_path
def cleanup_text_embedding(self):
if self.prompt_embeds is not None:
# we are caching on disk, don't save in memory
self.prompt_embeds = None
if self.dop_prompt_embeds is not None:
self.dop_prompt_embeds = None
def load_prompt_embedding(self, device=None):
if not self.is_text_embedding_cached:
@@ -2131,6 +2164,13 @@ class TextEmbeddingFileItemDTOMixin:
if self.prompt_embeds is None:
# load it from disk
self.prompt_embeds = PromptEmbeds.load(self.get_text_embedding_path())
if self.dataset_config.diff_output_preservation and self.dop_prompt_embeds is None:
dop_path = self.get_dop_text_embedding_path()
if dop_path == self.get_text_embedding_path():
# no trigger word in caption, same embedding
self.dop_prompt_embeds = self.prompt_embeds
else:
self.dop_prompt_embeds = PromptEmbeds.load(dop_path)
class TextEmbeddingCachingMixin:
def __init__(self: 'AiToolkitDataset', **kwargs):
@@ -2152,13 +2192,21 @@ class TextEmbeddingCachingMixin:
file_item.latent_load_device = self.sd.device
text_embedding_path = file_item.get_text_embedding_path(recalculate=True)
# (path, caption) pairs to encode for this item
encode_targets = [(text_embedding_path, file_item.caption)]
if self.dataset_config.diff_output_preservation:
dop_path = file_item.get_dop_text_embedding_path(recalculate=True)
if dop_path != text_embedding_path:
# trigger word was in the caption, cache the DOP version too
encode_targets.append((dop_path, file_item.caption_dop))
# only process if not saved to disk
if not os.path.exists(text_embedding_path):
encode_targets = [t for t in encode_targets if not os.path.exists(t[0])]
if len(encode_targets) > 0:
# load if not loaded
if not did_move:
self.sd.set_device_state_preset('cache_text_encoder')
did_move = True
if file_item.encode_control_in_text_embeddings and file_item.control_path is not None:
ctrl_img_list = []
control_path_list = file_item.control_path
@@ -2185,7 +2233,10 @@ class TextEmbeddingCachingMixin:
ctrl_img = ctrl_img_list[0]
else:
ctrl_img = ctrl_img_list
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(file_item.caption, control_images=ctrl_img)
for path, caption in encode_targets:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(caption, control_images=ctrl_img)
prompt_embeds.save(path)
del prompt_embeds
elif (
getattr(self.sd, 'encode_first_frame_in_text_embeddings', False)
and self.dataset_config.do_i2v
@@ -2205,13 +2256,16 @@ class TextEmbeddingCachingMixin:
)
if self.sd.has_multiple_control_images:
ctrl_img = [ctrl_img]
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(file_item.caption, control_images=ctrl_img)
for path, caption in encode_targets:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(caption, control_images=ctrl_img)
prompt_embeds.save(path)
del prompt_embeds
file_item.tensor = None
else:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(file_item.caption)
# save it
prompt_embeds.save(text_embedding_path)
del prompt_embeds
for path, caption in encode_targets:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(caption)
prompt_embeds.save(path)
del prompt_embeds
file_item.is_text_embedding_cached = True
i += 1
# restore device state