Initial Switch to Lark parser
This commit is contained in:
@@ -1,79 +0,0 @@
|
||||
from lark import Lark, Tree, Token, Transformer
|
||||
|
||||
grammar = """
|
||||
?start: item+
|
||||
|
||||
item: embedding
|
||||
| WORD
|
||||
| function
|
||||
| QUOTED_STRING
|
||||
|
||||
value: WORD
|
||||
| embedding
|
||||
| function
|
||||
| QUOTED_STRING
|
||||
|
||||
function: sum_function
|
||||
| neg_function
|
||||
|
||||
sum_function: "sum(" value ("|" value)* ")"
|
||||
neg_function: "neg(" value ")"
|
||||
|
||||
embedding: "embedding:" WORD
|
||||
|
||||
|
||||
|
||||
WORD: /[A-Za-z0-9_-]+/
|
||||
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
|
||||
|
||||
%import common.WS
|
||||
%ignore WS
|
||||
"""
|
||||
|
||||
def flatten_tree(tree):
|
||||
if isinstance(tree, Token):
|
||||
return [str(tree)]
|
||||
else:
|
||||
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
||||
|
||||
# Initialize the parser
|
||||
parser = Lark(grammar, start='start', parser='earley')
|
||||
|
||||
sample_texts = [
|
||||
# "A cat embedding:sdfds sum(rabbit|sum(embedding:sdfds|rose|neg(car))) cat",
|
||||
"A cat sum(rabbit|sum(embedding:sdfds|rose|neg(car))) cat",
|
||||
# "embedding:scj0_jg [<cat:+embedding:scj0_jg>:<water:+mountain>:0.5] is the embedding:scj0_jg at a <dog:+[embedding:scj0_jg cat:embedding:scj0_jg:0.1]-<water:+ocean>>",
|
||||
]
|
||||
|
||||
class MyTransformer(Transformer):
|
||||
# def WORD(self, items):
|
||||
# return items
|
||||
|
||||
def item(self, items):
|
||||
print(items)
|
||||
return items
|
||||
|
||||
# def value(self, items):
|
||||
# for item in items:
|
||||
# print(item)
|
||||
# return items
|
||||
|
||||
# def embedding(self, items):
|
||||
# return items
|
||||
#
|
||||
# def sum_function(self, items):
|
||||
# return items
|
||||
#
|
||||
# def neg_function(self, items):
|
||||
# return items
|
||||
|
||||
|
||||
# Parse the sample text
|
||||
for sample_action_text in sample_texts:
|
||||
parsed_tree = parser.parse(sample_action_text)
|
||||
MyTransformer().transform(parsed_tree)
|
||||
# print(parsed_tree.pretty())
|
||||
print(flatten_tree(parsed_tree))
|
||||
|
||||
# Flatten the tree into a list of tokens
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from lark import Lark
|
||||
|
||||
from .grammar import grammar
|
||||
|
||||
PromptParser = Lark(grammar, start="start", parser="earley")
|
||||
@@ -0,0 +1,24 @@
|
||||
grammar = """
|
||||
?start: item+
|
||||
|
||||
item: embedding
|
||||
| WORD
|
||||
| function
|
||||
| QUOTED_STRING
|
||||
|
||||
function: sum_function
|
||||
| neg_function
|
||||
| norm_function
|
||||
|
||||
sum_function: "sum(" item* ("|" item)* ")"
|
||||
neg_function: "neg(" item* ")"
|
||||
norm_function: "norm(" item* ")"
|
||||
|
||||
embedding: "embedding:" WORD
|
||||
|
||||
WORD: /[A-Za-z0-9,_-]+/
|
||||
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
|
||||
|
||||
%import common.WS
|
||||
%ignore WS
|
||||
"""
|
||||
@@ -0,0 +1,54 @@
|
||||
from lark import Transformer, Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action, build_prompt_segment
|
||||
from custom_nodes.ClipStuff.lib.actions.neg import NegAction
|
||||
from custom_nodes.ClipStuff.lib.actions.norm import NormAction
|
||||
from custom_nodes.ClipStuff.lib.actions.sum import SumAction
|
||||
from custom_nodes.ClipStuff.lib.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
class PromptTransformer(Transformer):
|
||||
# def WORD(self, items):
|
||||
# return items
|
||||
|
||||
def __init__(self, tokenizer: SD1Tokenizer):
|
||||
super().__init__()
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def item(self, items: list[Token]):
|
||||
for item in items:
|
||||
if isinstance(item, Action):
|
||||
return item
|
||||
|
||||
if isinstance(item, PromptSegment):
|
||||
return item
|
||||
|
||||
if item.type == "WORD":
|
||||
return build_prompt_segment(str(item), self.tokenizer)
|
||||
elif item.type == "QUOTED_STRING":
|
||||
# Remove the quotes
|
||||
unquoted = item[1:-1]
|
||||
# Replace escaped quotes with quotes
|
||||
unescaped = unquoted.replace("\\\"", "\"").replace("\\\'", "\'")
|
||||
return build_prompt_segment(unescaped, self.tokenizer)
|
||||
elif item.type == "embedding":
|
||||
return build_prompt_segment(item, self.tokenizer)
|
||||
elif item.type == "function":
|
||||
return item
|
||||
else:
|
||||
raise Exception("Unknown item type: " + str(item.type))
|
||||
|
||||
def embedding(self, items):
|
||||
return build_prompt_segment(f'{self.tokenizer.embedding_identifier}{items[0]}', self.tokenizer)
|
||||
|
||||
def function(self, items):
|
||||
for item in items:
|
||||
if item.data == 'sum_function':
|
||||
return SumAction(item.children[0], item.children[1:])
|
||||
elif item.data == 'neg_function':
|
||||
return NegAction(item.children)
|
||||
elif item.data == 'norm_function':
|
||||
return NormAction(item.children)
|
||||
else:
|
||||
raise Exception("Unknown function type: " + str(item.data))
|
||||
@@ -0,0 +1,8 @@
|
||||
from lark import Token
|
||||
|
||||
|
||||
def flatten_tree(tree):
|
||||
if isinstance(tree, Token):
|
||||
return [str(tree)]
|
||||
else:
|
||||
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
||||
+8
-103
@@ -1,94 +1,12 @@
|
||||
import re
|
||||
from typing import Union
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions.all_actions import (
|
||||
NudgeAction,
|
||||
ArithAction,
|
||||
ALL_START_CHARS,
|
||||
ALL_END_CHARS,
|
||||
ALL_ACTIONS,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.base import (
|
||||
Action,
|
||||
build_prompt_segment,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
|
||||
is_any_action_segment,
|
||||
is_action_segment,
|
||||
)
|
||||
|
||||
from custom_nodes.ClipStuff.lib.parser import PromptParser
|
||||
from custom_nodes.ClipStuff.lib.parser.transformer import PromptTransformer
|
||||
from custom_nodes.ClipStuff.lib.prompt_segment import PromptSegment
|
||||
|
||||
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
|
||||
|
||||
# TODO: Get embedding identifier from tokenizer
|
||||
tokenizer_regex = re.compile(
|
||||
fr"""
|
||||
\d+\.\d+ # Capture decimals
|
||||
|
|
||||
(?:(?!embedding:)[\w\s]|embedding:[a-zA-Z0-9_]+)+ # Capture sequences of characters, including "embedding:"
|
||||
|
|
||||
\d+ # Capture whole numbers
|
||||
|
|
||||
[:+-{re.escape("".join(ALL_START_CHARS))}{re.escape("".join(ALL_END_CHARS))}] # Capture special characters including start and end characters
|
||||
""",
|
||||
re.VERBOSE
|
||||
)
|
||||
def tokenize(text: str) -> list[str]:
|
||||
# Captures:
|
||||
# 1. Words
|
||||
# 2. Numbers(1.0, 1)
|
||||
# 3. Special characters(ALL_START_CHARS, ALL_END_CHARS, :, +, -)
|
||||
tokens = re.findall(tokenizer_regex, text)
|
||||
print(tokens)
|
||||
return [token.strip() for token in tokens]
|
||||
|
||||
|
||||
|
||||
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, tokenizer)
|
||||
# If we get here, it's a text segment
|
||||
return build_prompt_segment(tokens.pop(0), tokenizer)
|
||||
|
||||
def parse(tokens: list[str], tokenizer: SD1Tokenizer) -> list[PromptSegment | Action]:
|
||||
parsed = []
|
||||
while tokens:
|
||||
if tokens[0] == '':
|
||||
tokens.pop(0)
|
||||
continue
|
||||
print("Parse: Checking token: " + tokens[0])
|
||||
if tokens[0] in ALL_START_CHARS:
|
||||
parsed.append(parse_segment(tokens, tokenizer))
|
||||
else:
|
||||
parsed.append(build_prompt_segment(tokens.pop(0), tokenizer))
|
||||
return parsed
|
||||
|
||||
|
||||
def parse_special_tokens(string) -> list[str]:
|
||||
out = []
|
||||
current = ""
|
||||
|
||||
for char in string:
|
||||
if char in ALL_START_CHARS:
|
||||
out += [current]
|
||||
current = char
|
||||
elif char in ALL_END_CHARS:
|
||||
out += [current + char]
|
||||
current = ""
|
||||
else:
|
||||
current += char
|
||||
out += [current]
|
||||
return out
|
||||
|
||||
|
||||
def parse_segment_actions(string, tokenizer: SD1Tokenizer) -> list[SegOrAction]:
|
||||
tokens = tokenize(string)
|
||||
parsed = parse(tokens, tokenizer)
|
||||
return parsed
|
||||
|
||||
class TokenDict:
|
||||
def __init__(self,
|
||||
token_id: int,
|
||||
@@ -124,28 +42,15 @@ class MyTokenizer(SD1Tokenizer):
|
||||
else:
|
||||
pad_token = 0
|
||||
|
||||
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
|
||||
for segment in parsed_actions:
|
||||
if isinstance(segment, Action):
|
||||
print(segment.depth_repr())
|
||||
else:
|
||||
print(segment.depth_repr())
|
||||
parsed_prompt = PromptParser.parse(text)
|
||||
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
|
||||
|
||||
# reshape token array to CLIP input size
|
||||
batched_segments = []
|
||||
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
||||
# batched_segments.append(batch)
|
||||
batch_size = 1
|
||||
for segment in parsed_actions:
|
||||
for segment in parsed_actions.children:
|
||||
num_tokens = segment.token_length()
|
||||
# determine if we're going to try and keep the tokens in a single batch
|
||||
is_large = num_tokens >= self.max_word_length
|
||||
@@ -171,7 +76,7 @@ class MyTokenizer(SD1Tokenizer):
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
||||
batched_segments.append(batch)
|
||||
|
||||
for batch in batched_segments:
|
||||
batch_size_info(batch)
|
||||
# for batch in batched_segments:
|
||||
# batch_size_info(batch)
|
||||
|
||||
return batched_segments
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
lark
|
||||
Reference in New Issue
Block a user