diff --git a/toolkit/advanced_prompt_embeds.py b/toolkit/advanced_prompt_embeds.py new file mode 100644 index 00000000..61ebdf3b --- /dev/null +++ b/toolkit/advanced_prompt_embeds.py @@ -0,0 +1,164 @@ +from safetensors.torch import load_file, save_file + + +class AdvancedPromptEmbeds: + """ + Flexible container for prompt embedding tensors. + + Each value passed in must be a list of tensors, where each item in the + list corresponds to a single item in the batch (list length == batch size). + Do not store more than one tensor per batch item under the same key — if + you need multiple tensors per batch item, give them different key names. + + Usage: + pe = AdvancedPromptEmbeds( + prompt_embeds=[t0, t1, t2], # one tensor per batch item + pooled_embeds=[p0, p1, p2], + ) + + pe.prompt_embeds # -> [t0, t1, t2] + pe['prompt_embeds'] # -> [t0, t1, t2] + pe.keys() # -> ['prompt_embeds', 'pooled_embeds'] + + # add more after init + pe.extra = [e0, e1, e2] + pe['extra2'] = [e0, e1, e2] + pe.set('extra3', [e0, e1, e2]) + pe.update(extra4=[e0, e1, e2]) + """ + + def __init__(self, **kwargs): + self._store = {} + for key, value in kwargs.items(): + if not isinstance(value, list): + value = [value] + self._store[key] = value + + def __getattr__(self, name): + if name.startswith("_"): + raise AttributeError(name) + store = self.__dict__.get("_store", {}) + if name in store: + return store[name] + raise AttributeError(f"{type(self).__name__!s} has no attribute {name!r}") + + def __setattr__(self, name, value): + if name.startswith("_"): + super().__setattr__(name, value) + else: + if not isinstance(value, list): + value = [value] + self._store[name] = value + + def set(self, key, value): + if not isinstance(value, list): + value = [value] + self._store[key] = value + + def update(self, **kwargs): + for key, value in kwargs.items(): + if not isinstance(value, list): + value = [value] + self._store[key] = value + + def keys(self): + return list(self._store.keys()) + + def __getitem__(self, key): + return self._store[key] + + def __setitem__(self, key, value): + if not isinstance(value, list): + value = [value] + self._store[key] = value + + def __contains__(self, key): + return key in self._store + + def to(self, *args, **kwargs): + new_pe = AdvancedPromptEmbeds() + for key, value in self._store.items(): + new_pe._store[key] = [v.to(*args, **kwargs) for v in value] + return new_pe + + def detach(self): + new_pe = AdvancedPromptEmbeds() + for key, value in self._store.items(): + new_pe._store[key] = [v.detach() for v in value] + return new_pe + + def clone(self): + new_pe = AdvancedPromptEmbeds() + for key, value in self._store.items(): + new_pe._store[key] = [v.clone() for v in value] + return new_pe + + def expand_to_batch(self, batch_size): + new_pe = AdvancedPromptEmbeds() + for key, value in self._store.items(): + if len(value) == 1: + new_pe._store[key] = value * batch_size + elif len(value) == batch_size: + new_pe._store[key] = value + else: + raise ValueError( + f"Cannot expand key {key!r}: expected list of length 1 or {batch_size}, got {len(value)}" + ) + return new_pe + + def save(self, path): + data = {} + metadata = {"class_name": self.__class__.__name__} + for key, value in self._store.items(): + if len(value) != 1: + raise ValueError( + f"Cannot save key {key!r}: expected list of length 1, got {len(value)}" + ) + data[key] = value[0] + save_file(data, path, metadata=metadata) + + @classmethod + def load(cls, path=None): + if path is not None: + loaded = load_file(path) + else: + raise ValueError("Must provide a path") + + metadata = loaded.metadata() + if metadata.get("class_name") != cls.__name__: + raise ValueError( + f"Metadata class_name {metadata.get('class_name')!r} does not match expected {cls.__name__!r}" + ) + + data = {} + for key in loaded.keys(): + data[key] = loaded[key] + + return cls(**data) + + @classmethod + def concat_prompt_embeds( + cls, prompt_embeds: list["AdvancedPromptEmbeds"], padding_side: str = "right" + ): + embeds = {} + for pe in prompt_embeds: + for key in pe.keys(): + if key not in embeds: + embeds[key] = [] + embeds[key].append(pe[key]) + return cls(**embeds) + + @classmethod + def split_prompt_embeds(cls, concatenated: "AdvancedPromptEmbeds", num_parts=None): + if num_parts is None: + # use length of first item as num_parts + num_parts = len(concatenated[concatenated.keys()[0]]) + split_embeds = [cls() for _ in range(num_parts)] + for key in concatenated.keys(): + values = concatenated[key] + if len(values) != num_parts: + raise ValueError( + f"Cannot split key {key!r}: expected list of length {num_parts}, got {len(values)}" + ) + for i in range(num_parts): + split_embeds[i]._store[key] = values[i] diff --git a/toolkit/prompt_utils.py b/toolkit/prompt_utils.py index ef3a7da8..6628fd61 100644 --- a/toolkit/prompt_utils.py +++ b/toolkit/prompt_utils.py @@ -8,6 +8,8 @@ import random from toolkit.train_tools import get_torch_dtype import itertools +from safetensors import safe_open +from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds if TYPE_CHECKING: from toolkit.config_modules import SliderTargetConfig @@ -145,6 +147,12 @@ class PromptEmbeds: :param path: The path to load the prompt embeds from. :return: An instance of PromptEmbeds. """ + # first check if it is advanced prompt embed file + f = safe_open(path, framework='pt') + metadata = f.metadata() + if metadata is not None and metadata.get("class_name", "") == "AdvancedPromptEmbeds": + return AdvancedPromptEmbeds.load(path=path) + state_dict = load_file(path, device='cpu') text_embeds = [] pooled_embeds = None @@ -245,6 +253,9 @@ class EncodedPromptPair: def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"], padding_side: str = "right") -> PromptEmbeds: + # check if first item has a classmethod of concat_prompt_embeds + if hasattr(prompt_embeds[0].__class__, "concat_prompt_embeds"): + return prompt_embeds[0].__class__.concat_prompt_embeds(prompt_embeds, padding_side=padding_side) # --- pad text_embeds --- if isinstance(prompt_embeds[0].text_embeds, (list, tuple)): text_embeds = [] @@ -335,6 +346,8 @@ def concat_prompt_pairs(prompt_pairs: list[EncodedPromptPair]): def split_prompt_embeds(concatenated: PromptEmbeds, num_parts=None) -> List[PromptEmbeds]: + if hasattr(concatenated.__class__, "split_prompt_embeds"): + return concatenated.__class__.split_prompt_embeds(concatenated, num_parts=num_parts) if num_parts is None: # use batch size num_parts = concatenated.text_embeds.shape[0]