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:
@@ -1,5 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
|
from safetensors import safe_open
|
||||||
from safetensors.torch import load_file, save_file
|
from safetensors.torch import load_file, save_file
|
||||||
|
|
||||||
|
|
||||||
@@ -137,6 +138,8 @@ class AdvancedPromptEmbeds:
|
|||||||
def save(self, path):
|
def save(self, path):
|
||||||
data = {}
|
data = {}
|
||||||
metadata = {"class_name": self.__class__.__name__}
|
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():
|
for key, value in self._store.items():
|
||||||
if len(value) != 1:
|
if len(value) != 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -150,6 +153,8 @@ class AdvancedPromptEmbeds:
|
|||||||
def load(cls, path=None):
|
def load(cls, path=None):
|
||||||
if path is not None:
|
if path is not None:
|
||||||
loaded = load_file(path)
|
loaded = load_file(path)
|
||||||
|
with safe_open(path, framework="pt") as f:
|
||||||
|
metadata = f.metadata()
|
||||||
else:
|
else:
|
||||||
raise ValueError("Must provide a path")
|
raise ValueError("Must provide a path")
|
||||||
|
|
||||||
@@ -157,7 +162,18 @@ class AdvancedPromptEmbeds:
|
|||||||
for key in loaded.keys():
|
for key in loaded.keys():
|
||||||
data[key] = loaded[key]
|
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
|
@classmethod
|
||||||
def concat_prompt_embeds(
|
def concat_prompt_embeds(
|
||||||
|
|||||||
@@ -2092,7 +2092,10 @@ class TextEmbeddingFileItemDTOMixin:
|
|||||||
self._dop_text_embedding_path: Union[str, None] = None
|
self._dop_text_embedding_path: Union[str, None] = None
|
||||||
# blank caption embeds used for caption dropout when caching text embeddings
|
# blank caption embeds used for caption dropout when caching text embeddings
|
||||||
self._blank_text_embedding_path: Union[str, None] = None
|
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._loaded_text_embedding_path: Union[str, None] = None
|
||||||
|
self._caption_was_dropped = False
|
||||||
self.is_text_embedding_cached = False
|
self.is_text_embedding_cached = False
|
||||||
self.text_embedding_load_device = 'cpu'
|
self.text_embedding_load_device = 'cpu'
|
||||||
self.text_embedding_version = 1
|
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 inject_trigger_into_prompt('', trigger=self.trigger_word, add_if_not_present=True)
|
||||||
return ''
|
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):
|
def get_blank_text_embedding_path(self: 'FileItemDTO', recalculate=False):
|
||||||
if self._blank_text_embedding_path is not None and not recalculate:
|
if self._blank_text_embedding_path is not None and not recalculate:
|
||||||
return self._blank_text_embedding_path
|
return self._blank_text_embedding_path
|
||||||
@@ -2185,17 +2210,24 @@ class TextEmbeddingFileItemDTOMixin:
|
|||||||
return
|
return
|
||||||
if self.prompt_embeds is None:
|
if self.prompt_embeds is None:
|
||||||
te_path = self.get_text_embedding_path()
|
te_path = self.get_text_embedding_path()
|
||||||
|
self._caption_was_dropped = False
|
||||||
if self.dataset_config.caption_dropout_rate > 0:
|
if self.dataset_config.caption_dropout_rate > 0:
|
||||||
# get a random float form 0 to 1
|
# get a random float form 0 to 1
|
||||||
rand = random.random()
|
rand = random.random()
|
||||||
if rand < self.dataset_config.caption_dropout_rate:
|
if rand < self.dataset_config.caption_dropout_rate:
|
||||||
# drop the caption by using the cached blank embedding
|
# drop the caption by using the cached blank embedding
|
||||||
te_path = self.get_blank_text_embedding_path()
|
te_path = self.get_blank_text_embedding_path()
|
||||||
|
self._caption_was_dropped = True
|
||||||
# load it from disk
|
# load it from disk
|
||||||
self.prompt_embeds = PromptEmbeds.load(te_path)
|
self.prompt_embeds = PromptEmbeds.load(te_path)
|
||||||
self._loaded_text_embedding_path = te_path
|
self._loaded_text_embedding_path = te_path
|
||||||
if self.dataset_config.diff_output_preservation and self.dop_prompt_embeds is None:
|
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:
|
if dop_path == self._loaded_text_embedding_path:
|
||||||
# no trigger word in caption, same embedding
|
# no trigger word in caption, same embedding
|
||||||
self.dop_prompt_embeds = self.prompt_embeds
|
self.dop_prompt_embeds = self.prompt_embeds
|
||||||
@@ -2234,6 +2266,11 @@ class TextEmbeddingCachingMixin:
|
|||||||
if blank_path != text_embedding_path:
|
if blank_path != text_embedding_path:
|
||||||
# cache the dropout caption embedding (blank, or trigger word only)
|
# cache the dropout caption embedding (blank, or trigger word only)
|
||||||
encode_targets.append((blank_path, file_item.get_dropout_caption()))
|
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
|
# only process if not saved to disk
|
||||||
encode_targets = [t for t in encode_targets if not os.path.exists(t[0])]
|
encode_targets = [t for t in encode_targets if not os.path.exists(t[0])]
|
||||||
if len(encode_targets) > 0:
|
if len(encode_targets) > 0:
|
||||||
|
|||||||
Reference in New Issue
Block a user