Initial Switch to Lark parser

This commit is contained in:
Michael Poutre
2023-08-25 00:06:17 -07:00
parent 07cec3f68f
commit 6b8a1ea243
7 changed files with 100 additions and 182 deletions
-79
View File
@@ -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
+5
View File
@@ -0,0 +1,5 @@
from lark import Lark
from .grammar import grammar
PromptParser = Lark(grammar, start="start", parser="earley")
+24
View File
@@ -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
"""
+54
View File
@@ -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))
+8
View File
@@ -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
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
lark