From f5bd4cb91dc45a1fa151dae6771d736c5be8c32c Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Fri, 3 Nov 2023 01:33:54 +0200 Subject: [PATCH] Completely new wildcard parser --- __init__.py | 3 - utils/default_functions.txt | 38 ++++++ utils/other.py | 50 -------- utils/parse.py | 212 ++++++++++++++++++++++++++++++++ utils/wildcards.py | 239 ++++++++++++++++++++++++++++++++---- 5 files changed, 462 insertions(+), 80 deletions(-) create mode 100644 utils/default_functions.txt delete mode 100644 utils/other.py create mode 100644 utils/parse.py diff --git a/__init__.py b/__init__.py index c70849a..a949346 100644 --- a/__init__.py +++ b/__init__.py @@ -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, } diff --git a/utils/default_functions.txt b/utils/default_functions.txt new file mode 100644 index 0000000..0da1e26 --- /dev/null +++ b/utils/default_functions.txt @@ -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 = >; $y = $eval($e - $start) + $ascseq($x; x; 0; $y; $step; $step) +} + +def $coollora($lora; $start; $step) { + var $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 => +} diff --git a/utils/other.py b/utils/other.py deleted file mode 100644 index 9bcb886..0000000 --- a/utils/other.py +++ /dev/null @@ -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,) diff --git a/utils/parse.py b/utils/parse.py new file mode 100644 index 0000000..05baba1 --- /dev/null +++ b/utils/parse.py @@ -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()) diff --git a/utils/wildcards.py b/utils/wildcards.py index b38644e..a18ff09 100644 --- a/utils/wildcards.py +++ b/utils/wildcards.py @@ -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(?Pj?)\s+(?P\$[a-z]+)\((?P(\s*\$[a-z]+\s*;)*(\s*\$[a-z]+\s*)?)\)\s*\{(?P.*?)(?\$[a-z]+)\((?P.*?)(?(\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\$[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[A-Za-z0-9_/.-]+)(\+(?P[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("")