diff --git a/lib/action/base.py b/lib/action/base.py index eb2a952..b3794d6 100644 --- a/lib/action/base.py +++ b/lib/action/base.py @@ -4,6 +4,7 @@ from typing import Union, List, TypedDict, Tuple from torch import Tensor from torch.nn import Embedding +from transformers.models.clip.modeling_clip import CLIPTextTransformer from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment @@ -18,6 +19,7 @@ class PostModifiers(TypedDict): A dictionary of post modifiers for an action result. """ position_embed_scale: Union[float, None] + bypass_pos_embed: Union[bool, None] class Action(ABC): @@ -82,6 +84,13 @@ class Action(ABC): """ pass + def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None: + """ + For actions that need access to the TextTransformer, this method is called. Results are expected to be returned via get_result still. + :param transformer: An instance of CLIPTextTransformer + """ + pass + def depth_repr(self, depth: int = 1) -> str: raise NotImplementedError() @@ -97,6 +106,11 @@ class SingleArgAction(Action, ABC): segments.append(seg_or_action) return segments + def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None: + for seg_or_action in self.arg: + if isinstance(seg_or_action, Action): + seg_or_action.process_with_transformer(transformer, embedding_module) + def __init__(self, arg: List[Union[PromptSegment, Action]]): # TODO: Target is a list now... what does this mean for us.. self.arg = arg @@ -117,6 +131,12 @@ class MultiArgAction(Action, ABC): return segments + def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None: + for arg in self.all_args: + for seg_or_action in arg: + if isinstance(seg_or_action, Action): + seg_or_action.process_with_transformer(transformer, embedding_module) + def __init__( self, args: List[List[Union[PromptSegment, Action]]], diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py index d48628b..b96f267 100644 --- a/lib/fun_clip_stuff.py +++ b/lib/fun_clip_stuff.py @@ -1,4 +1,4 @@ -from typing import Optional, Tuple, Union, List, TypedDict +from typing import Optional, Tuple, Union, List, TypedDict, TYPE_CHECKING from importlib.metadata import version as import_version from packaging import version @@ -14,6 +14,11 @@ from transformers.models.clip.modeling_clip import ( from custom_nodes.KepPromptLang.lib.action.base import Action from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction +if TYPE_CHECKING: + from torch import Tensor + from custom_nodes.KepPromptLang.lib.action.base import PostModifiers + + def slerp(val, low, high): low = low.unsqueeze(0) high = high.unsqueeze(0) @@ -46,11 +51,10 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, ) -> torch.Tensor: - if input_dicts is None: raise ValueError("You have to specify input_dicts") - batches = [] + batches: List[List[Tensor | Tuple[Tensor, PostModifiers] | Action]] = [] pos_modifiers: List[List[PosModifier]] = [] for batch_idx, batch in enumerate(input_dicts): results = [] @@ -58,15 +62,23 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): token_idx = 0 for seg_or_action in batch: if isinstance(seg_or_action, Action): - action_result = seg_or_action.get_result(self.token_embedding) + action_result: Union[ + Tensor, Tuple[Tensor, PostModifiers] + ] = seg_or_action.get_result(self.token_embedding) if isinstance(action_result, tuple): result, post_modifiers = action_result - if post_modifiers["position_embed_scale"] is not None: + if post_modifiers.get("position_embed_scale", None) is not None: post_modifiers["start_idx"] = token_idx post_modifiers["end_idx"] = ( token_idx + seg_or_action.token_length() ) - batch_pos_modifiers.append(post_modifiers) + + if post_modifiers.get("bypass_pos_embed", False): + post_modifiers["start_idx"] = token_idx + post_modifiers["end_idx"] = ( + token_idx + seg_or_action.token_length() + ) + batch_pos_modifiers.append(post_modifiers) else: result = action_result else: @@ -88,14 +100,26 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): else: embeds.append(torch.cat(batch, dim=-2)) + # Iterate over the batches and apply the pos modifiers to the position embeddings then add them to the embeddings for idx, batch_pos_modifiers in enumerate(pos_modifiers): position_embeddings = self.position_embedding(position_ids) if len(batch_pos_modifiers) > 0: print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}") + # Apply each pos modifier to the position embeddings at the specified indices for post_modifier in batch_pos_modifiers: - position_embeddings[ - 0, post_modifier["start_idx"] : post_modifier["end_idx"] - ] *= post_modifier["position_embed_scale"] + if post_modifier.get("bypass_pos_embed", False): + position_embeddings[ + 0, post_modifier["start_idx"] : post_modifier["end_idx"] + ] = 0 + elif post_modifier["position_embed_scale"] is not None: + position_embeddings[ + 0, post_modifier["start_idx"] : post_modifier["end_idx"] + ] *= post_modifier["position_embed_scale"] + else: + raise ValueError( + "Pos modifier must have a scale or bypass_pos_embed" + ) + # Add the possibly modified position embeddings to the embeddings embeds[idx] = embeds[idx] + position_embeddings embeddings = torch.cat(embeds, dim=0)