feat(Action): Add experimental pooledAvg, pooler actions(_exp- prefix)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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?"
|
||||
)
|
||||
@@ -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?"
|
||||
)
|
||||
Reference in New Issue
Block a user