diff --git a/lib/actions/arith.py b/lib/actions/arith.py index 0a577fd..cb813e5 100644 --- a/lib/actions/arith.py +++ b/lib/actions/arith.py @@ -1,5 +1,6 @@ from typing import Callable, Union +from comfy.sd1_clip import SD1Tokenizer from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment @@ -7,7 +8,7 @@ class ArithAction(Action): START_CHAR = "<" END_CHAR = ">" - def __init__(self, base_segment: PromptSegment, ops: dict[str, list[Union[str, Action]]]): + def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]): self.base_segment = base_segment self.ops = ops @@ -19,8 +20,11 @@ class ArithAction(Action): if isinstance(self.base_segment, Action): base_segment_repr = self.base_segment.depth_repr(depth + 1) out += "\t" * depth + f"base_segment={base_segment_repr}\n" + elif isinstance(self.base_segment, PromptSegment): + out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}' else: out += "\t" * depth + f'base_segment="{self.base_segment}",' + for op_key, ops in self.ops.items(): for op in ops: out += "\n" + "\t" * depth + f'"{op_key}":[\n' @@ -28,7 +32,7 @@ class ArithAction(Action): op_repr = op.depth_repr(depth + 2) out += "\t" * (depth + 1) + f"{op_repr}\n" else: - out += "\t" * (depth + 1) + f'"{op}",\n' + out += "\t" * (depth + 1) + f'{op.depth_repr()},\n' out += "\t" * depth + "]," out += "\n" + "\t" * (depth - 1) + ")" return out @@ -40,7 +44,8 @@ class ArithAction(Action): tokens: list[str], start_chars: list[str], end_chars: list[str], - parent_parser: Callable[[list[str]], Union[str, 'Action']], + parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']], + tokenizer: SD1Tokenizer, ) -> Action: """ Parse an arithmetic action from a list of tokens @@ -57,7 +62,7 @@ class ArithAction(Action): assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token # Parse base segment - base_segment = parent_parser(tokens) + base_segment = parent_parser(tokens, tokenizer) token = tokens.pop(0) assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token @@ -67,7 +72,7 @@ class ArithAction(Action): while tokens[0] != cls.END_CHAR: op_char = tokens.pop(0) assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char - ops[op_char].append(parent_parser(tokens)) + ops[op_char].append(parent_parser(tokens, tokenizer)) token = tokens.pop(0) assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token diff --git a/lib/actions/base.py b/lib/actions/base.py index 8d2a373..fd67d89 100644 --- a/lib/actions/base.py +++ b/lib/actions/base.py @@ -1,6 +1,10 @@ from abc import ABC, abstractmethod from typing import Callable, Union +from torch import Tensor + +from comfy.sd1_clip import SD1Tokenizer + class Action(ABC): @property @@ -20,11 +24,45 @@ class Action(ABC): tokens: list[str], start_chars: list[str], end_chars: list[str], - parent_parser: Callable[[list[str]], Union[str, 'Action']], + parent_parser: Callable[[list[str], SD1Tokenizer], Union[str, 'Action']], + tokenizer: SD1Tokenizer, ) -> 'Action': pass def depth_repr(self, depth=1): raise NotImplementedError() -PromptSegment = str | Action +class PromptSegment: + def __init__(self, text, tokenizer: SD1Tokenizer): + self.text = text + self.tokens: list[Union[int, Tensor]] = [] + self.process_text(tokenizer) + + def depth_repr(self, depth=1): + out = f'"{self.text}"(' + + cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens)) + out += ", ".join(cleaned_tokens) + + out += ")" + return out + + def process_text(self, tokenizer: SD1Tokenizer): + split_text = self.text.split(" ") + for word in split_text: + if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None: + embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n') + + get_embed_ret = tokenizer._try_get_embedding(embedding_name) + embedding = get_embed_ret[0] + leftover = get_embed_ret[1] + if embedding is None: + print(f"warning, embedding:{embedding_name} does not exist, ignoring") + else: + self.tokens.append(embedding) + + if leftover != "": + word = leftover + else: + continue + self.tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1]) diff --git a/lib/actions/nudge.py b/lib/actions/nudge.py index 182a6c9..c906ce2 100644 --- a/lib/actions/nudge.py +++ b/lib/actions/nudge.py @@ -1,6 +1,7 @@ from typing import Optional, Union, Callable -from custom_nodes.ClipStuff.lib.actions.base import Action +from comfy.sd1_clip import SD1Tokenizer +from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment class NudgeAction(Action): @@ -9,8 +10,8 @@ class NudgeAction(Action): def __init__( self, - base_segment: Union[str, Action], - target: Union[str, Action], + base_segment: PromptSegment | Action, + target: Union[PromptSegment, Action], weight: Optional[float] = None, ): self.base_segment = base_segment @@ -26,13 +27,13 @@ class NudgeAction(Action): base_segment_repr = self.base_segment.depth_repr(depth + 1) out += "\t" * depth + f"base_segment={base_segment_repr}\n" else: - out += "\t" * depth + f'base_segment="{self.base_segment}",\n' + out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n' if isinstance(self.target, Action): target_repr = self.target.depth_repr(depth + 1) out += "\t" * depth + f"target={target_repr},\n" else: - out += "\t" * depth + f"target={self.target},\n" + out += "\t" * depth + f"target={self.target.depth_repr()},\n" out += "\t" * depth + f"weight={self.weight},\n" out += "\t" * (depth - 1) + ")" return out @@ -43,7 +44,8 @@ class NudgeAction(Action): tokens: list[str], start_chars: list[str], end_chars: list[str], - parent_parser: Callable[[list[str]], Union[str, Action]], + parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action], + tokenizer: SD1Tokenizer, ) -> Action: """ Parse a nudge action from a list of tokens @@ -62,13 +64,13 @@ class NudgeAction(Action): assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token # Parse base segment - base_segment = parent_parser(tokens) + base_segment = parent_parser(tokens, tokenizer) token = tokens.pop(0) assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token # Parse target segment - target_segment = parent_parser(tokens) + target_segment = parent_parser(tokens, tokenizer) # Parse weight if it exists weight = None diff --git a/lib/tokenizer.py b/lib/tokenizer.py index a9085b1..6636651 100644 --- a/lib/tokenizer.py +++ b/lib/tokenizer.py @@ -9,7 +9,7 @@ from custom_nodes.ClipStuff.lib.actions import ( ALL_END_CHARS, ALL_ACTIONS, ) -from custom_nodes.ClipStuff.lib.actions.base import Action +from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment from custom_nodes.ClipStuff.lib.actions.lib import ( is_any_action_segment, is_action_segment, @@ -32,19 +32,22 @@ def tokenize(text: str) -> list[str]: -def parse_segment(tokens: list[str]) -> Union[str, Action]: +def parse_segment(tokens: list[str], tokenizer: SD1Tokenizer) -> PromptSegment | Action: + print("Parse segment: Checking token: " + tokens[0]) for action in ALL_ACTIONS: if tokens[0] == action.START_CHAR: - return action.parse_segment(tokens, ALL_START_CHARS, ALL_END_CHARS, parse_segment) - # If we get here, it's a word - return tokens.pop(0) -def parse(tokens: list[str]) -> list[Union[str, Action]]: + return action.parse_segment(tokens, ALL_START_CHARS, ALL_END_CHARS, parse_segment, tokenizer) + # If we get here, it's a text segment + return PromptSegment(tokens.pop(0), tokenizer) + +def parse(tokens: list[str], tokenizer: SD1Tokenizer) -> list[PromptSegment | Action]: parsed = [] while tokens: + print("Parse: Checking token: " + tokens[0]) if tokens[0] in ALL_START_CHARS: - parsed.append(parse_segment(tokens)) + parsed.append(parse_segment(tokens, tokenizer)) else: - parsed.append(tokens.pop(0)) + parsed.append(PromptSegment(tokens.pop(0), tokenizer)) return parsed @@ -65,9 +68,9 @@ def parse_special_tokens(string) -> list[str]: return out -def parse_segment_actions(string) -> list[Union[str, NudgeAction, ArithAction]]: +def parse_segment_actions(string, tokenizer: SD1Tokenizer) -> list[PromptSegment | NudgeAction | ArithAction]: tokens = tokenize(string) - parsed = parse(tokens) + parsed = parse(tokens, tokenizer) return parsed class TokenDict: @@ -103,72 +106,23 @@ class MyTokenizer(SD1Tokenizer): else: pad_token = 0 - parsed_actions = parse_segment_actions(text) + parsed_actions = parse_segment_actions(text, self) - nudge_start = kwargs.get("nudge_start") - nudge_end = kwargs.get("nudge_end") - - if nudge_start is not None and nudge_end is not None: - nudge_start = int(nudge_start) - nudge_end = int(nudge_end) - - # tokenize words - tokens: list[list[Action | int ]] = [] - - for action in parsed_actions: - nudge_weight = None - nudge_to_id = None - arith_ops = None - if isinstance(action, str): - segment_to_tokenize = action - elif isinstance(action, NudgeAction): - segment_to_tokenize = 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): - segment_to_tokenize = 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]] + # nudge_start = kwargs.get("nudge_start") + # nudge_end = kwargs.get("nudge_end") + # + # if nudge_start is not None and nudge_end is not None: + # nudge_start = int(nudge_start) + # nudge_end = int(nudge_end) + # + # # tokenize words + for segment in parsed_actions: + if isinstance(segment, Action): + print(segment.depth_repr()) else: - raise Exception(f"Unexpected action type: {type(action)}") + print(segment.depth_repr()) - to_tokenize = segment_to_tokenize.split(' ') - to_tokenize = [x for x in to_tokenize if x != ""] - - for word in to_tokenize: - # if we find an embedding, deal with the embedding - if word.startswith(self.embedding_identifier) and self.embedding_directory is not None: - embedding_name = word[len(self.embedding_identifier):].strip('\n') - embed, leftover = self._try_get_embedding(embedding_name) - if embed is None: - print(f"warning, embedding:{embedding_name} does not exist, ignoring") - else: - if len(embed.shape) == 1: - tokens.append([TokenDict(token_id=embed)]) - else: - tokens.append([ - TokenDict(token_id=embed[x]) - for x in range(embed.shape[0]) - ]) - # if we accidentally have leftover text, continue parsing using leftover, else move on to next word - if leftover != "": - word = leftover - else: - continue - # parse word - tokens.append([t] for t in self.tokenizer(word)["input_ids"][1:-1]) - tokens.append([TokenDict( - token_id=t, - nudge_id=nudge_to_id, - nudge_weight=nudge_weight, - nudge_start=nudge_start, - nudge_end=nudge_end, - arith_ops=arith_ops - ) for t in self.tokenizer(word)["input_ids"][1:-1]]) + tokens: list[list[Action | int ]] = [] # reshape token array to CLIP input size batched_tokens = []