Switch to using PromptSegment to represent text chunks also add tokens
This commit is contained in:
+10
-5
@@ -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
@@ -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
@@ -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
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user