diff --git a/lib/__init__.py b/lib/__init__.py index c332b50..1e66cb0 100644 --- a/lib/__init__.py +++ b/lib/__init__.py @@ -3,6 +3,9 @@ from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction from custom_nodes.KepPromptLang.lib.actions.neg import NegAction from custom_nodes.KepPromptLang.lib.actions.norm import NormAction + +from custom_nodes.KepPromptLang.lib.actions.pooled_avg import PooledAvgAction +from custom_nodes.KepPromptLang.lib.actions.pooler import PoolerAction from custom_nodes.KepPromptLang.lib.actions.pos_scale import PosScaleAction from custom_nodes.KepPromptLang.lib.actions.post_pos import PostPosAction from custom_nodes.KepPromptLang.lib.actions.rand import RandAction @@ -23,4 +26,6 @@ register_action(AverageAction) register_action(ScaleDims) register_action(SetDims) register_action(PosScaleAction) +register_action(PoolerAction) register_action(PostPosAction) +register_action(PooledAvgAction) diff --git a/lib/actions/action_utils.py b/lib/actions/action_utils.py index 3f676f2..f23f556 100644 --- a/lib/actions/action_utils.py +++ b/lib/actions/action_utils.py @@ -1,3 +1,5 @@ +from typing import List + from torch import Tensor from torch.nn import Embedding @@ -9,3 +11,10 @@ def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Te if isinstance(seg_or_action, Action): return seg_or_action.get_result(embedding_module) return seg_or_action.get_embeddings(embedding_module) + +def get_total_length(args: List[SegOrAction]) -> int: + total_length = 0 + for seg_or_action in args: + total_length += seg_or_action.token_length() + + return total_length diff --git a/lib/actions/pooled_avg.py b/lib/actions/pooled_avg.py new file mode 100644 index 0000000..d68af44 --- /dev/null +++ b/lib/actions/pooled_avg.py @@ -0,0 +1,61 @@ +import torch +from torch.nn import Embedding +from transformers.modeling_outputs import BaseModelOutputWithPooling +from transformers.models.clip.modeling_clip import CLIPTextTransformer + +from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction +from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length +from custom_nodes.KepPromptLang.lib.fun_clip_stuff import ( + PromptLangCLIPTextEmbeddings, + PrompLangCLIPTextTransformer, +) +from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment + + +class PooledAvgAction(SingleArgAction): + grammar = 'pooledAvg(" arg+ ")"' + name = "_exp-pooledAvg" + chars = ["[", "]"] + + def __init__(self, args): + super().__init__(args) + self.result = None + + def token_length(self) -> int: + """ + PooledAvg returns the average of the last hidden state, so the length is 1 + """ + return 1 + + def process_with_transformer( + self, transformer: CLIPTextTransformer, embedding_module: Embedding + ) -> None: + """ """ + # SOT + tokens + EOT + eot_token = embedding_module.num_embeddings - 1 + print("Using EOT token", eot_token) + #TODO: Play with impact of padding on pooled output + arg_length = get_total_length(self.arg) + + # SOT + arg length + EOT + empty_tokens = [[49406] + [eot_token] * (arg_length + 1)] + + transformer_results: BaseModelOutputWithPooling = transformer( + [ + [PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])], + [PromptSegment(text="[SOT]", tokens=[49406])] + + self.arg + + [PromptSegment(text="[EOT]", tokens=[eot_token])], + ] + ) + self.result = ( + transformer_results.last_hidden_state[1, 1:-1, :].mean(dim=0).unsqueeze(0).unsqueeze(0) + ) + + def get_result(self, embedding_module: Embedding) -> torch.Tensor: + if self.result is not None: + return self.result + + raise Exception( + "PooledAvg action result is not set. Did you forget to call process_with_transformer?" + ) diff --git a/lib/actions/pooler.py b/lib/actions/pooler.py new file mode 100644 index 0000000..91ae99a --- /dev/null +++ b/lib/actions/pooler.py @@ -0,0 +1,59 @@ +import torch +from torch.nn import Embedding +from transformers.modeling_outputs import BaseModelOutputWithPooling +from transformers.models.clip.modeling_clip import CLIPTextTransformer + +from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction +from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length +from custom_nodes.KepPromptLang.lib.fun_clip_stuff import ( + PromptLangCLIPTextEmbeddings, + PrompLangCLIPTextTransformer, +) +from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment + + +class PoolerAction(SingleArgAction): + grammar = 'pooler(" arg+ ")"' + name = "_exp-pooler" + chars = ["[", "]"] + + def __init__(self, args): + super().__init__(args) + self.result = None + + def token_length(self) -> int: + """ + Pooler returns the embedding of the EOT token, so the length is 1 + """ + return 1 + + def process_with_transformer( + self, transformer: CLIPTextTransformer, embedding_module: Embedding + ) -> None: + """ """ + # SOT + tokens + EOT + eot_token = embedding_module.num_embeddings - 1 + print("Using EOT token", eot_token) + + arg_length = get_total_length(self.arg) + + # SOT + arg length + EOT + empty_tokens = [[49406] + [eot_token] * (arg_length + 1)] + + transformer_results: BaseModelOutputWithPooling = transformer( + [ + [PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])], + [PromptSegment(text="[SOT]", tokens=[49406])] + + self.arg + + [PromptSegment(text="[EOT]", tokens=[eot_token])], + ] + ) + self.result = transformer_results.pooler_output[1].unsqueeze(0).unsqueeze(0) + + def get_result(self, embedding_module: Embedding) -> torch.Tensor: + if self.result is not None: + return self.result + + raise Exception( + "Pooled action result is not set. Did you forget to call process_with_transformer?" + )