Things might work now

This commit is contained in:
asagi4
2023-12-22 01:30:39 +02:00
parent 96671d55c8
commit 8a2d6810da
3 changed files with 69 additions and 39 deletions
+28 -18
View File
@@ -1,33 +1,43 @@
$comment($anything) = {}
$descseq($expr, $var=x, $startval=1, $endval=0, $startat=0, $step=0.1) = {
[SEQ<%- for step in steps(0, $startval - $endval, $step) -%>
<%- set $var "=" round($startval - step, 2) -%>
<%- if $var ">=" $endval -%>
:$expr:<= round($startat + step, 2) =>
$comment($doc="Ignores its parameters") = {}
$debugwc=
$valueof($x, $doc="Does nothing, but documents the default value of a function parameter") = { "$valueof("$x")" }
$default($expr, $defval, $doc="returns $defval if $expr is $valueof(something)") = {
$('$defval' if '$expr'.startswith("'$valueof'") else '$expr')
}
$genericseq($expr, $var, $startval, $endval, $startat, $step, $increment, $op, $hidden=true) = {
[SEQ<%- for step in steps(0, $(round(abs(($endval) - ($startval)), 2)), $step) -%>
<%- set $var "=" round($startval + loop.index0 * ($increment), 2) -%>
<%- if $var $op ($endval) -%>
:$expr:"<=" round(($startat) + step, 2) "=>"
<%- endif -%>
<%- endfor -%>
<%- set $var "=" $endval -%>:$expr:1]
}
$ascseq($expr, $var=x, $startval=0, $endval=1, $startat=0, $step=0.1) = {
[SEQ<%- for step in steps(0, round($endval - $startval, 2), $step) -%>
<%- set $var "=" round($startval+step, 2) -%>
<%- if $var < $endval -%>
:$expr:<= round($startat + step, 2) =>
<%- endif -%>
<%- endfor -%>
<%- set $var "=" $endval -%>:$expr:1]
$ascseq($expr, $var=x, $startval=0, $endval=1, $startat=0, $step=0.1, $increment=$valueof(step)) = {
$i = $default($increment, $step)
$genericseq($expr, $var, $startval, $endval, $startat, $step, $i, $op="<")
}
$warmlora($lora, $start=0, $end=1.0, $step=0.1, $startat=0) = {
$descseq($expr, $var=x, $startval=1, $endval=0, $startat=$valueof(step), $step=0.1, $increment=$valueof(step)) = {
$i = $default($increment, $step)
$s = $default($startat, $step)
$genericseq($expr, $var, $startval, $endval, $s, $step, $increment=-($i), $op=">")
}
$warmlora($lora, $start=0, $end=1.0, $step=0.1, $startat=0, $increment=$valueof(step)) = {
$i = $default($increment, $step)
$x = <lora:$lora:<= x + $start "=>>"
$y = $($end - $start)
$ascseq($x, x, 0, $y, $startat, $step)
$ascseq($x, x, 0, $y, $startat, $step, $i)
}
$coollora($lora, $start, $step=0.1, $startat=0) = {
$coollora($lora, $start, $step=0.1, $startat=0, $decrement=$valueof(step)) = {
$i = $default($decrement, $step)
$x = <lora:$lora:<= x =>>
$descseq($x, x, $start, 0, $startat, $step)
$descseq($x, x, $start, 0, $startat, $step, $i)
}
$rectmask($x, $y, $size) = {
+31 -19
View File
@@ -65,8 +65,6 @@ def const(x):
class Context:
POISON = object()
def __init__(self):
self.vars = ChainMap()
@@ -84,9 +82,6 @@ class Context:
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) == "$":
@@ -116,15 +111,22 @@ def showarg(a):
def print_context_functions(ctx):
log.info("Functions available:")
log.info("- $help(), shows this help")
log.info("- $($expr), evaluates $expr as jinja")
def p(x):
print("MUWildcard help:", x)
p("Functions available:")
p("- $help(), shows this help")
p("- $($expr), evaluates $expr as jinja")
for v, val in ctx.vars.items():
if isinstance(val, tuple):
log.info(f"- ${v}({','.join(showarg(a) for a in val[1])})")
if isinstance(val, tuple) and "hidden" not in [a[0] for a in val[1]]:
p(f"- ${v}({','.join(showarg(a) for a in val[1])})")
MAGIC_FUNCTIONS = {"$": eval, "help": print_context_functions}
def debug(ctx, x):
print("MUWildCard Debug:", x)
MAGIC_FUNCTIONS = {"$": eval, "help": print_context_functions, "debug": debug}
class TestVisitor(Interpreter):
@@ -140,10 +142,17 @@ class TestVisitor(Interpreter):
var = varname(var)
args = self.visit_children(argspec)
found_defval = None
for x, defval in args:
if found_defval and not defval:
raise TypeError(f"Invalid function definition for {var}, must define default for {x}")
found_defval = defval
args = []
with self.ctx as c:
for arg in argspec.children:
x, defval = self.visit(arg)
if found_defval and not defval:
raise TypeError(f"Invalid function definition for {var}, must define default for {x}")
found_defval = defval
args.append((x, defval))
# Allow $fn($a=1, $b=$a)
if x not in c.vars:
c.set(x, defval)
with self.ctx as locals:
res = (locals, args, function_body)
self.ctx.set(var, res)
@@ -257,9 +266,12 @@ class TestVisitor(Interpreter):
return final_prompt, self.ctx
parser = lark.Lark(definition, parser="earley")
def parse(x, ctx=None):
try:
r = TestVisitor(ctx).visit(raw_parse(x))
r = TestVisitor(ctx).visit(parser.parse(x))
if r is None:
return x, None
return r
@@ -268,10 +280,10 @@ def parse(x, ctx=None):
return x, None
def raw_parse(text, p="earley", **kwargs):
x = lark.Lark(definition, parser=p, debug=True, **kwargs)
def debug_parse(text, p="earley", **kwargs):
x = lark.Lark(definition, parser=p, debug=True)
return x.parse(text)
def pparse(text, p="earley", **kwargs):
print(raw_parse(text, p, **kwargs).pretty())
print(debug_parse(text, p, **kwargs).pretty())
+10 -2
View File
@@ -84,6 +84,9 @@ def replace_lora_tags(text, seed):
return text
global_ctx = None
def wildcard_prompt_handler(json_data):
for node_id in json_data["prompt"].keys():
if json_data["prompt"][node_id]["class_type"] == CLASS_NAME:
@@ -92,12 +95,14 @@ def wildcard_prompt_handler(json_data):
def handle_wildcard_node(json_data, node_id):
global global_ctx
wildcard_info = json_data.get("extra_data", {}).get("extra_pnginfo", {}).get(CLASS_NAME, {})
n = json_data["prompt"][node_id]
if not (n["inputs"].get("use_pnginfo") and node_id in wildcard_info):
text = MUSimpleWildcard.select(n["inputs"]["text"], n["inputs"]["seed"])
ctx = read_preamble()
text, _ = parse(text, ctx)
if not global_ctx or "$debugwc" in text:
global_ctx = read_preamble()
text, _ = parse(text, global_ctx)
text = replace_lora_tags(text, n["inputs"]["seed"])
if text.strip() != n["inputs"]["text"].strip():
@@ -111,6 +116,7 @@ def read_preamble():
curfile = Path(__file__)
defaults = curfile.parent / "default_functions.txt"
with open(defaults, "r") as f:
log.info("Reading functions from %s", defaults)
return parse(f.read())[1]
@@ -208,4 +214,6 @@ except ImportError:
print("Could not install wildcard prompt handler, node won't work")
if __name__ == "__main__":
print("Start parse")
ctx = read_preamble()
print("End parse")