feat: Add support for arithmetic ops

This commit is contained in:
Michael Poutre
2023-08-19 18:44:42 -07:00
parent 088400f3fa
commit 265d5c3557
2 changed files with 157 additions and 38 deletions
+8
View File
@@ -58,6 +58,14 @@ class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
nudge_start = 0
nudge_end = 768
inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
elif token[0].arith_ops is not None:
for op, id_list in token[0].arith_ops.items():
if op == '+':
for this_id in id_list:
inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
elif op == '-':
for this_id in id_list:
inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
position_embeddings = self.position_embedding(position_ids)
embeddings = inputs_embeds + position_embeddings
+149 -38
View File
@@ -1,4 +1,5 @@
from typing import Union, TypedDict, Optional
from enum import Enum
from typing import Union, TypedDict, Optional, Literal
from comfy.sd1_clip import SD1Tokenizer, escape_important, unescape_important, parse_parentheses
@@ -23,48 +24,142 @@ def token_weights(string, current_weight):
out += [(x, current_weight)]
return out
def parse_brackets(string):
class SpecialChars(str, Enum):
nudge_start = '['
nudge_end = ']'
arith_start = '<'
arith_end = '>'
START_CHARS = [SpecialChars.nudge_start, SpecialChars.arith_start]
END_CHARS = [SpecialChars.nudge_end, SpecialChars.arith_end]
def parse_special_tokens(string):
out = []
current = ""
for char in string:
if char == '[':
if char in START_CHARS:
out += [current]
current = "["
elif char == ']':
out += [current + ']']
current = char
elif char in END_CHARS:
out += [current + char]
current = ""
else:
current += char
out += [current]
return out
def parse_nudges(string) -> list[tuple[str, Union[str, None]]]:
out = []
for nudge_segment in parse_brackets(string):
if nudge_segment == "":
# class TokenActions:
# def __init__(self, is_nudge=None, is_arith=None, nudge_weight=None, base_segment=None):
# if is_nudge is None:
# self.is_nudge = False
# else:
# self.is_nudge = is_nudge
#
# if is_arith is None:
# self.is_arith = False
# else:
# self.is_arith = is_arith
#
# self.nudge_weight = nudge_weight
#
#
#
# def validate(self):
# if self.is_arith and self.is_nudge:
# raise Exception("Token action cannot be both arith and nudge")
#
# if self.is_nudge and self.bas
class NudgeAction:
def __init__(self, base_segment=None, weight: Optional[float] = None, target=None):
self.base_segment = base_segment
self.weight = weight
self.target = target
class ArithAction:
def __init__(self, base_segment: str, ops_str: str):
self.base_segment = base_segment
self.ops = self.process_ops_string(ops_str)
@classmethod
def process_ops_string(cls, ops_string):
supported_ops = ['+', '-']
# dict[[Union[Literal['add'], Literal['subtract']]], str]
ops_dict = {'+': [], '-': []}
buff = ''
curr_op_char = ''
for char in ops_string:
if char in supported_ops:
# We have a buffer
if buff != '':
# Add op string
ops_dict[curr_op_char] += [buff]
# Reset buffer
buff = ''
# Set new current op char
curr_op_char = char
continue
else:
# No buffer, the start of processing
curr_op_char = char
else:
# Append char to buffer
buff += char
# Add last op to dict
ops_dict[curr_op_char] += [buff]
return ops_dict
def parse_token_actions(string) -> list[Union[str, NudgeAction, ArithAction]]:
out: list[Union[str, NudgeAction, ArithAction]] = []
for prompt_segment in parse_special_tokens(string):
if prompt_segment == "":
continue
if nudge_segment[0] != '[' and nudge_segment[-1] != ']':
out += [(nudge_segment, None, None)]
if prompt_segment[0] not in START_CHARS and prompt_segment[-1] not in END_CHARS:
out += [prompt_segment]
continue
nudge_segment = nudge_segment[1:-1]
sep_idx = nudge_segment.find(":")
if sep_idx < 0:
out += [(nudge_segment, None, None)]
is_nudge = is_arith = False
if prompt_segment[0] == SpecialChars.nudge_start and prompt_segment[-1] == SpecialChars.nudge_end:
is_nudge = True
elif prompt_segment[0] == SpecialChars.arith_start and prompt_segment[-1] == SpecialChars.arith_end:
is_arith = True
prompt_segment = prompt_segment[1:-1]
word_sep_idx = prompt_segment.find(":")
# No word seperator, add whole segment
if word_sep_idx < 0:
out += [prompt_segment]
continue
nudge_to = nudge_segment[sep_idx+1:]
weight = None
base_segment = prompt_segment[:word_sep_idx]
weight_sep_idx = nudge_to.find(":")
if weight_sep_idx >= 0:
[nudge_to, weight] = nudge_to.split(":")
weight = float(weight)
if is_nudge:
trailing_segment = prompt_segment[word_sep_idx + 1:]
weight_sep_idx = trailing_segment.find(":")
# Has a weight(base_word:nudge_to:1.4)
if weight_sep_idx >= 0:
[nudge_to, weight] = trailing_segment.split(":")
weight = float(weight)
else:
# No weight(base_word:trailing_segment)
nudge_to = trailing_segment
weight = None
out += [NudgeAction(base_segment, weight, nudge_to)]
elif is_arith:
arith_op_string = prompt_segment[word_sep_idx + 1:]
out += [ArithAction(base_segment, arith_op_string)]
out += [(nudge_segment[:sep_idx], nudge_to, weight)]
return out
# class TokenDict(TypedDict):
# token_id: int
# weight: float
@@ -73,8 +168,11 @@ def parse_nudges(string) -> list[tuple[str, Union[str, None]]]:
#
class TokenDict:
def __init__(self, token_id: int, weight: float = None, nudge_id=None, nudge_weight=None, nudge_start: int = None,
nudge_end: int = None):
def __init__(self,
token_id: int,
weight: float = None,
nudge_id=None, nudge_weight=None, nudge_start: int = None, nudge_end: int = None,
arith_ops: dict[str, list[str]] = None):
if weight is None:
self.weight = 1.0
else:
@@ -86,6 +184,8 @@ class TokenDict:
self.nudge_index_start = nudge_start
self.nudge_index_stop = nudge_end
self.arith_ops = arith_ops
class MyTokenizer(SD1Tokenizer):
@@ -101,7 +201,7 @@ class MyTokenizer(SD1Tokenizer):
else:
pad_token = 0
parsed_nudges = parse_nudges(text)
parsed_actions = parse_token_actions(text)
nudge_start = None
nudge_end = None
@@ -113,19 +213,29 @@ class MyTokenizer(SD1Tokenizer):
#tokenize words
tokens: list[list[TokenDict]] = []
for token_segment, nudge_to_token, nudge_weight in parsed_nudges:
for action in parsed_actions:
nudge_weight = None
nudge_to_id = None
arith_ops = None
if isinstance(action, str):
token_segment = action
elif isinstance(action, NudgeAction):
token_segment = action.base_segment
nudge_to_id = self.tokenizer(action.target)["input_ids"][1:-1][0]
nudge_weight = action.weight
if nudge_weight is None:
nudge_weight = 0.5
elif isinstance(action, ArithAction):
token_segment = action.base_segment
arith_ops = action.ops
for op in arith_ops:
arith_ops[op] = [self.tokenizer(word)["input_ids"][1:-1][0] for word in arith_ops[op]]
else:
raise Exception(f"Unexpected action type: {type(action)}")
to_tokenize = token_segment.split(' ')
to_tokenize = [x for x in to_tokenize if x != ""]
# if token_segment == ' ':
# continue
if nudge_weight is None:
nudge_weight = .5
nudge_to_id = None
if nudge_to_token is not None:
# self.convert_tokens_to_ids
nudge_to_id = self.tokenizer(nudge_to_token)["input_ids"][1:-1][0]
for word in to_tokenize:
#if we find an embedding, deal with the embedding
@@ -153,7 +263,8 @@ class MyTokenizer(SD1Tokenizer):
nudge_id=nudge_to_id,
nudge_weight=nudge_weight,
nudge_start=nudge_start,
nudge_end=nudge_end
nudge_end=nudge_end,
arith_ops=arith_ops
) for t in self.tokenizer(word)["input_ids"][1:-1]])
#reshape token array to CLIP input size