From 265d5c35575dbcf8a113f3966c42ea1f657bee8a Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Sat, 19 Aug 2023 18:44:42 -0700 Subject: [PATCH] feat: Add support for arithmetic ops --- lib/fun_clip_stuff.py | 8 ++ lib/tokenizer.py | 187 +++++++++++++++++++++++++++++++++--------- 2 files changed, 157 insertions(+), 38 deletions(-) diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py index f7bee23..40743c4 100644 --- a/lib/fun_clip_stuff.py +++ b/lib/fun_clip_stuff.py @@ -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 diff --git a/lib/tokenizer.py b/lib/tokenizer.py index 58c5c8d..0240c4d 100644 --- a/lib/tokenizer.py +++ b/lib/tokenizer.py @@ -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