Switch to using PromptSegment to represent text chunks also add tokens

This commit is contained in:
Michael Poutre
2023-08-25 00:06:17 -07:00
parent 235a681f70
commit 21eecda007
4 changed files with 87 additions and 88 deletions
+10 -5
View File
@@ -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
+40 -2
View File
@@ -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])
+10 -8
View File
@@ -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
+27 -73
View File
@@ -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 = []