feat: Add support for arithmetic ops
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user