Completely new wildcard parser
This commit is contained in:
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 =>
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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("")
|
||||
|
||||
Reference in New Issue
Block a user