From c3d45e306ee35ff83a2ee6d0c70adc2b9b5ab9d8 Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Sun, 29 Oct 2023 18:28:59 -0700 Subject: [PATCH] feat: Add support for post embedding modifiers --- lib/action/base.py | 11 +++++++++-- lib/fun_clip_stuff.py | 45 ++++++++++++++++++++++++++++++++++++++----- 2 files changed, 49 insertions(+), 7 deletions(-) diff --git a/lib/action/base.py b/lib/action/base.py index f6fb227..eb2a952 100644 --- a/lib/action/base.py +++ b/lib/action/base.py @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from enum import Enum -from typing import Union, List +from typing import Union, List, TypedDict, Tuple from torch import Tensor from torch.nn import Embedding @@ -13,6 +13,13 @@ class ActionArity(Enum): SINGLE = 1 MULTI = 2 +class PostModifiers(TypedDict): + """ + A dictionary of post modifiers for an action result. + """ + position_embed_scale: Union[float, None] + + class Action(ABC): @property @abstractmethod @@ -67,7 +74,7 @@ class Action(ABC): pass @abstractmethod - def get_result(self, embedding_module: Embedding) -> Tensor: + def get_result(self, embedding_module: Embedding) -> Union[Tensor, Tuple[Tensor, PostModifiers]]: """ Get the result of this action. This is called when the embeddings are being calculated. :param embedding_module: The embedding module to use to get the base embeddings for tokens. diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py index 5fe6af4..aa2a8bc 100644 --- a/lib/fun_clip_stuff.py +++ b/lib/fun_clip_stuff.py @@ -1,4 +1,4 @@ -from typing import Optional, Tuple, Union, List +from typing import Optional, Tuple, Union, List, TypedDict import torch from transformers import CLIPTextConfig @@ -23,6 +23,17 @@ def slerp(val, low, high): res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high return res + +class PosModifier(TypedDict): + """ + A dictionary of post modifiers for an action result. + """ + + position_embed_scale: Union[float] + start_idx: Union[int] + end_idx: Union[int] + + class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): def __init__(self, config: CLIPTextConfig): super().__init__(config) @@ -39,14 +50,30 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): raise ValueError("You have to specify input_dicts") batches = [] + pos_modifiers: List[List[PosModifier]] = [] for batch_idx, batch in enumerate(input_dicts): results = [] + batch_pos_modifiers = [] + token_idx = 0 for seg_or_action in batch: if isinstance(seg_or_action, Action): - results.append(seg_or_action.get_result(self.token_embedding)) + action_result = 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: + 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: - results.append(seg_or_action.get_embeddings(self.token_embedding)) + result = seg_or_action.get_embeddings(self.token_embedding) + results.append(result) + token_idx += seg_or_action.token_length() batches.append(results) + pos_modifiers.append(batch_pos_modifiers) seq_length = batches[0][0].shape[-2] @@ -60,8 +87,16 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings): else: embeds.append(torch.cat(batch, dim=-2)) - position_embeddings = self.position_embedding(position_ids) - embeddings = torch.cat(embeds, dim=0) + position_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}") + for post_modifier in batch_pos_modifiers: + position_embeddings[ + 0, post_modifier["start_idx"] : post_modifier["end_idx"] + ] *= post_modifier["position_embed_scale"] + embeds[idx] = embeds[idx] + position_embeddings + embeddings = torch.cat(embeds, dim=0) return embeddings