Completely new wildcard parser

This commit is contained in:
asagi4
2023-11-22 01:02:07 +02:00
parent 69d4eadfab
commit f5bd4cb91d
5 changed files with 462 additions and 80 deletions
-3
View File
@@ -1,10 +1,7 @@
from .utils import jinja_render as j
from .utils import wildcards as w
from .utils import other as o
NODE_CLASS_MAPPINGS = {
"MUJinjaRender": j.MUJinjaRender,
"MUSimpleWildcard": w.MUSimpleWildcard,
"MUConditioningCutoff": o.MUConditioningCutoff,
"MUStringConcat": o.MUStringConcat,
}
+38
View File
@@ -0,0 +1,38 @@
def $descseq($expr; $var; $sw; $ew; $start; $step) {
[SEQ<%- for step in steps(0, $sw - $ew, $step) -%>
<%- set $var = round($sw - step, 2) -%>
<%- if $var >= $ew -%>
:$expr:<= round($start + step, 2) =>
<%- endif -%>
<%- endfor -%>
<%- set $var = $ew -%>:$expr:1]
}
def $ascseq($expr; $var; $sw; $ew; $start; $step) {
[SEQ<%- for step in steps(0, round($ew - $sw, 2), $step) -%>
<%- set $var = round($sw+step, 2) -%>
<%- if $var < $ew -%>
:$expr:<= round($start + step, 2) =>
<%- endif -%>
<%- endfor -%>
<%- set $var = $ew -%>:$expr:1]
}
def $warmlora($lora; $e; $start; $step) {
var $x = <lora:$lora:<= x + $start =>>; $y = $eval($e - $start)
$ascseq($x; x; 0; $y; $step; $step)
}
def $coollora($lora; $start; $step) {
var $x = <lora:$lora:<= x =>>
$descseq($x; x; $start; 0; 0; $step)
}
def $rectmask($x; $y; $size) {
MASK($x <= $x + $size =>, $y <= $y + $size =>)
}
# Evaluates expression as Jinja2 immediately
defj $eval($x) {
<= $x =>
}
-50
View File
@@ -1,50 +0,0 @@
class MUStringConcat:
@classmethod
def INPUT_TYPES(s):
t = ("STRING", {"default": ""})
return {
"optional": {
"string1": t,
"string2": t,
"string3": t,
"string4": t,
}
}
RETURN_TYPES = ("STRING",)
CATEGORY = "miscutils"
FUNCTION = "cat"
def cat(self, string1="", string2="", string3="", string4=""):
return string1 + string2 + string3 + string4
class MUConditioningCutoff:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"conds": ("CONDITIONING",),
"cutoff": ("FLOAT", {"min": 0.00, "max": 1.00, "default": 0.0, "step": 0.01}),
}
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "miscutils"
FUNCTION = "apply"
def apply(self, conds, cutoff):
res = []
new_start = 1.0
for c in conds:
end = c[1].get("end_percent", 0.0)
if 1.0 - end < cutoff:
continue
c = [c[0].clone(), c[1].copy()]
c[1]["start_percent"] = new_start
c[1]["end_percent"] = end
new_start = end
res.append(c)
return (res,)
+212
View File
@@ -0,0 +1,212 @@
import lark
import logging
logging.basicConfig(level=logging.DEBUG)
from lark.visitors import Interpreter, v_args
log = logging.getLogger("MUWildcard")
definition = r"""
%import common.NEWLINE
%import common.ESCAPED_STRING
%import common.WS_INLINE -> WS
quoted: ESCAPED_STRING
_WS: WS
start: (definition | expr )*
expr.0: prompt+
| NEWLINE
definition.10: var _WS? "=" _WS? expr? _TERM? -> var_definition
| var argument_spec _WS? "=" _WS? function_body -> function_definition
!block: "(" expr? ")"
| "{" expr? "}"
function_call.10: var argument_list | /\$/ argument_list
!prompt.0: quoted
| function_call
| var
| block
| STRING
| WS+ -> ws
| /[,;"]/
var.5: "$" "{" NAME "}" | "$" NAME
NAME: /[a-z]+/
argument_spec.10: "(" _WS? var? (_WS? _SEP _WS? var)* ")"
argument_list.10: "(" _WS? expr? (_WS? _SEP expr _WS?)*")"
function_body.20: "{" expr* "}"
_SEP.1: ","
_TERM.10: ";" | NEWLINE
STRING.0: /[^"$(){},;\n]+/
"""
from collections import ChainMap
from .jinja_render import render_jinja
def eval(x):
return render_jinja(f"<={x}=>")
MAGIC_FUNCTIONS = {"$": eval}
def const(x):
def f():
return x
return f
class Context:
POISON = object()
def __init__(self):
self.vars = ChainMap()
def __enter__(self):
self.vars = self.vars.new_child()
return self
def __exit__(self, *exc):
self.vars = self.vars.parents
return False
def set(self, name, value):
self.vars[str(name)] = value
def get(self, name, default=None):
return self.vars.get(str(name), default)
def poison(self, name):
self.vars[str(name)] = self.POISON
def varname(x):
if str(x) == "$":
return str(x)
return x.children[0].value
def flatten(seq):
if isinstance(seq, str):
yield seq
elif seq is None:
yield ""
else:
for x in seq:
yield from flatten(x)
def prompt(seq):
return "".join(flatten(seq))
def print_context_functions(ctx):
for v, val in ctx.vars.items():
print(val)
if isinstance(tuple, val):
print("function", v, val)
class TestVisitor(Interpreter):
def __init__(self, ctx=None):
super().__init__()
self.ctx = ctx or Context()
def __default__(self, tree):
return self.visit_children(tree)
@v_args(inline=True)
def function_definition(self, var, argspec, function_body):
var = varname(var)
args = [varname(a) for a in argspec.children]
with self.ctx as locals:
res = (locals, args, function_body)
self.ctx.set(var, res)
@v_args(inline=True)
def argument_list(self, *args):
return [prompt(self.visit_children(x)) for x in args]
@v_args(inline=True)
def quoted(self, value):
print("Quoted: '", value, "'")
value = value.replace('\\"', '"')
return value[1:][:-1]
@v_args(inline=True)
def function_call(self, var, arglist):
var = varname(var)
args = self.visit(arglist)
if var in MAGIC_FUNCTIONS:
return MAGIC_FUNCTIONS[var](*args)
try:
locals, params, function_body = self.ctx.get(var)
except TypeError:
raise TypeError(f"${var} is not a function")
print_context_functions(self.ctx)
with self.ctx as c:
if len(params) != len(args):
raise TypeError(f"Invalid number of arguments to function ${var}({','.join(f'${a}' for a in params)})")
for k, v in locals.vars.items():
c.set(k, v)
for a, v in zip(params, args):
c.set(a, const(v))
return prompt(self.visit(function_body)).strip()
@v_args(inline=True)
def var(self, name):
v = self.ctx.get(name)
if isinstance(v, tuple):
raise TypeError(f"${name} is a function, can't use as a variable")
if v:
return v()
else:
raise TypeError(f"${name} is undefined")
@v_args(inline=True)
def var_definition(self, var, definition=""):
name = var.children[0].value
if definition:
def resolve():
with self.ctx as c:
v = self.visit(definition)
c.set(name, const(v))
return v
else:
resolve = const("")
self.ctx.set(name, resolve)
def start(self, tree):
result = self.visit_children(tree)
final_prompt = "".join(flatten(result))
return final_prompt, self.ctx
def parse(x, ctx=None):
try:
return TestVisitor(ctx).visit(raw_parse(x))
except Exception as e:
log.error("Error: %s", e)
raise
def raw_parse(text, p="earley", **kwargs):
x = lark.Lark(definition, parser=p, debug=True, **kwargs)
return x.parse(text)
def pparse(text, p="earley", **kwargs):
print(raw_parse(text, p, **kwargs).pretty())
+212 -27
View File
@@ -3,15 +3,20 @@ from os import environ
import random
import re
from pathlib import Path
from server import PromptServer
from collections import ChainMap
from .jinja_render import render_jinja
from .parse import parse, pparse
# from .parse import parse
log = logging.getLogger("comfyui-misc-utils")
log.setLevel(logging.INFO)
CLASS_NAME = "MUSimpleWildcard"
def wildcard_prompt_handler(json_data):
log.info("Resolving wildcards...")
for node_id in json_data["prompt"].keys():
if json_data["prompt"][node_id]["class_type"] == CLASS_NAME:
handle_wildcard_node(json_data, node_id)
@@ -31,21 +36,176 @@ def handle_wildcard_node(json_data, node_id):
return json_data
PromptServer.instance.add_on_prompt_handler(wildcard_prompt_handler)
def find_and_remove(regexp, text, placeholder=""):
m = regexp.search(text)
res = {}
index = 0
def variable_substitution(text):
var_re = re.compile(r"(\$[a-z]+)\s*=([^;\n]*);?")
m = var_re.search(text)
while m:
var = m[1]
sub = m[2]
index += 1
ph = ""
if placeholder:
ph = "__" + placeholder + str(index) + "__"
res[ph] = m.groupdict()
else:
res[m["name"]] = m.groupdict()
s, e = m.span()
text = text[:s] + text[e:]
log.info("Substituting %s with '%s'", var, sub)
text = text.replace(var, sub)
m = var_re.search(text)
return text
text = text[:s] + ph + text[e:]
m = regexp.search(text)
return res, text
def function_definitions(text):
# $foo($a, $b...) { bodygoeshere \} }
func_re = re.compile(
r"def(?P<jinja>j?)\s+(?P<name>\$[a-z]+)\((?P<vars>(\s*\$[a-z]+\s*;)*(\s*\$[a-z]+\s*)?)\)\s*\{(?P<body>.*?)(?<!\\)}",
flags=re.MULTILINE | re.S,
)
found, text = find_and_remove(func_re, text)
res = {}
for k, v in found.items():
res[k] = {}
vars = [x.strip() for x in v["vars"].split(";") if x.strip()]
if len(vars) != len(set(vars)):
log.warning("Ignoring invalid function definition, duplicate vars: %s", vars)
continue
res[k]["vars"] = vars
res[k]["body"] = v["body"].strip()
res[k]["jinja"] = v["jinja"]
return res, text
def function_calls(text):
call_re = re.compile(r"(?P<name>\$[a-z]+)\((?P<args>.*?)(?<!\\)\)")
res, text = find_and_remove(call_re, text, placeholder="FUNC")
for k, v in res.items():
args = [x.strip().replace(r"\;", ";") for x in re.split(r"(?<!\\);", v["args"])]
res[k]["args"] = args
return res, text
def variable_definitions(text):
# name because find_and_remove expects t
var_re = re.compile(r"var\s+(?P<name>(\s*?\$[a-z]+\s*=[^;\n]*;?)+)")
defs, text = find_and_remove(var_re, text)
res = {}
for k, v in defs.items():
for d in (x for x in v["name"].strip().split(";") if x.strip()):
name, value = d.split("=", 1)
res[name.strip()] = value.strip()
return res, text
def push_context(ctx):
for k in ctx.keys():
ctx[k] = ctx[k].new_child()
def pop_context(ctx):
for k in ctx.keys():
ctx[k] = ctx[k].parents
def parent_context(ctx):
new_ctx = {}
for k in ctx.keys():
new_ctx[k] = ctx[k].parents
return new_ctx
class Nothing:
def __repr__(self):
return "Nothing"
RECURSIVE_VARIABLE = Nothing()
def read_preamble():
curfile = Path(__file__)
defaults = curfile.parent / "default_functions.txt"
with open(defaults, "r") as f:
defs = f.read()
_, defs = find_and_remove(re.compile("^#.*"), defs)
return defs
def read_preamble_new(func):
curfile = Path(__file__)
defaults = curfile.parent / "default_functions2.txt"
with open(defaults, "r") as f:
return func(f.read())
def init_context():
ctx = {"funcs": ChainMap(), "vars": ChainMap()}
curfile = Path(__file__)
defaults = curfile.parent / "default_functions.txt"
with open(defaults, "r") as f:
defs = f.read()
_, defs = find_and_remove(re.compile("^#.*"), defs)
funcs, _ = function_definitions(defs)
ctx["funcs"].update(funcs)
return ctx
def variable_substitution(text, ctx=None, allow_definitions=True, jinja_render=False):
if ctx is None:
ctx = init_context()
funcs, text = function_definitions(text)
if allow_definitions:
ctx["funcs"].update(funcs)
vars, text = variable_definitions(text)
ctx["vars"].update(vars)
calls, text = function_calls(text)
regex = re.compile(r"(?P<name>\$[a-z]+)(?!\()\b")
vars, text = find_and_remove(regex, text, placeholder="VARIABLE")
for placeholder, var in vars.items():
value = ctx["vars"].get(var["name"])
name = var["name"]
if value is None:
raise ValueError(f"Undefined variable {name}")
if value is RECURSIVE_VARIABLE:
raise ValueError(f"Recursive variable in definition of {name}")
push_context(ctx)
ctx["vars"][name] = RECURSIVE_VARIABLE
value = variable_substitution(value, ctx, allow_definitions=False)
pop_context(ctx)
text = text.replace(placeholder, value)
for k, v in calls.items():
f = ctx["funcs"].get(v["name"])
if not f:
text = text.replace(k, "")
log.warning("Function %s not found, ignoring", v["name"])
continue
if len(v["args"]) != len(f["vars"]):
log.warning("Invalid function call to %s, ignoring", v["name"])
ex = "; ".join(f["vars"])
log.warning("Call syntax: %s(%s)", v["name"], ex)
continue
push_context(ctx)
for var, val in zip(f["vars"], v["args"]):
r = variable_substitution(val, parent_context(ctx))
r = r.replace("\(", "(").replace("\)", ")")
ctx["vars"][var] = r
j = f.get("jinja")
replacement = variable_substitution(f["body"], ctx, jinja_render=j)
text = text.replace(k, replacement)
pop_context(ctx)
if jinja_render:
t = render_jinja(text)
if t.strip() != text.strip():
text = t
return text.strip()
class MUSimpleWildcard:
@@ -55,7 +215,7 @@ class MUSimpleWildcard:
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"default": ""}),
"text": ("STRING", {"default": "", "multiline": True}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
},
"optional": {"use_pnginfo": ("BOOLEAN", {"default": False})},
@@ -79,26 +239,51 @@ class MUSimpleWildcard:
with open(f, "r") as file:
return [l.strip() for l in file.readlines() if l.strip()]
except:
log.warning("Wildcard file not found for %s", name)
return [name]
@classmethod
def select(cls, text, seed):
cls.RAND.seed(seed)
matches = re.findall(r"(\$([A-Za-z0-9_/.-]+)(\+[0-9]+)?\$)", text)
for placeholder, wildcard, offset in matches:
wildcard_re = re.compile(r"\$(?P<name>[A-Za-z0-9_/.-]+)(\+(?P<offset>[0-9]+))?\$")
matches, text = find_and_remove(wildcard_re, text, placeholder="MU_WILDCARD")
for placeholder, value in matches.items():
state = None
ws = cls.read_wildcards(value["name"])
offset = int(value["offset"] or 0)
if offset:
offset = int(offset[1:])
cls.RAND.seed(seed + offset)
w = cls.RAND.choice(cls.read_wildcards(wildcard))
text = text.replace(placeholder, w, 1)
log.info("Selected wildcard %s for %s", w, placeholder)
if offset:
cls.RAND.seed(seed)
return text
# advance the state once and store it so that the next non-offset result stays deterministic
w = cls.RAND.choice(ws)
state = cls.RAND.getstate()
# advance the state until offset,
for _ in range(offset):
w = cls.RAND.choice(ws)
else:
w = cls.RAND.choice(ws)
text = text.replace(placeholder, w)
log.info("Selected wildcard %s for %s", w, value["name"])
if state:
cls.RAND.setstate(state)
return text.strip()
def doit(self, text, seed, extra_pnginfo, unique_id, use_pnginfo=False):
if use_pnginfo and unique_id in extra_pnginfo.get(CLASS_NAME, {}):
text = extra_pnginfo[CLASS_NAME][unique_id]
log.info("MUSimpleWildcard using prompt: %s", text)
text = variable_substitution(text)
return (text,)
newtext = variable_substitution(text)
if newtext != text:
log.info("MUSimpleWildcard result:\n%s", newtext)
return (newtext,)
try:
from server import PromptServer
PromptServer.instance.add_on_prompt_handler(wildcard_prompt_handler)
except ImportError:
print("Could not install wildcard prompt handler, node won't work")
if __name__ == "__main__":
_, ctx = read_preamble_new(parse)
pparse("")