Allow setting temporary variables in expressions

Improve sequence operator in expressions

Documentation improvements

Turn down eval debug spam a bit

Add ternary operator

Support static eval for constants in some more operator types
This commit is contained in:
blepping
2024-08-08 06:27:53 -06:00
parent 470b38231f
commit 2481e7263b
12 changed files with 156 additions and 95 deletions
+9
View File
@@ -4,6 +4,7 @@ Experimental and mathematically unsound (but fun!) sampling for [ComfyUI](https:
**Status**: In flux, may be useful but likely to change/break workflows frequently. Mainly for advanced users.
Feel free create a question in Discussions for usage help: [OCS Q&A Discussion](https://github.com/blepping/comfyui_overly_complicated_sampling/discussions/categories/q-a)
## Features
@@ -49,6 +50,14 @@ You may use filters and expressions in the text parameter input. See:
* [Filters](docs/filter.md)
* [Expressions](docs/expression.md)
## Integration
If you have [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) available, you will have access to many more blend and scaling modes as well as some extra features.
If you have [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) available, you will have access to many more noise types as well as the Power Filter feature.
If you're going to use OCS, I strongly recommend also installing those two node packs.
## Nodes
### `OCS Sampler`
+6
View File
@@ -21,6 +21,10 @@ Symbols (simple string type) are defined using `'symbol_name` - note the solitar
`;` can be used to sequence operations. I.E. `exp1 ; exp2` evaluates `exp1`, then `exp2` and then result of the expression is whatever `exp2` returned.
`:=` is used to assign to a temporary variable (see `set_var` below).
The expression language supports a C/JavaScript style ternary operator: `condition ? true_branch : false_branch` is the equivalent of `if(condition, true_branch, false_branch)`.
Like Python, a parenthesized expression with a trailing comma can be used to create an empty tuple. Example: `(1,)`
## Filter Variables
@@ -98,6 +102,8 @@ Available in model filters, with the exception of the `input` filter.
| <td colspan=3 align=left>Boolean negation</td> |
|⬤| `s_` | start:`I(null)`, end:`I(null)`, step:`I(null)` | `slice` |
| <td colspan=3 align=left>Creates a slice object from the `start`, `end`, `step` values. See Numpy [s_](https://numpy.org/doc/stable/reference/generatednumpy.s_.html)</td> |
|⬤| `set_var` | `SY`, `*` | `*` |
| <td colspan=3 align=left>Sets a temporary variable to the specified value and returns the value. Alias for the `:=` assignment operator. <br/> **Example**: `test1 := 2; set_var('test2, 10); test1 * test2`</td> |
|⬤| `unsafe_call` | `callable`, `*`\* | `*` |
| <td colspan=3 align=left>Allows calling an arbitrary callable. <br/> **Example:** `unsafe_call(some_callable, arg1, arg2, kwarg1 :> 123)`</td>
+2
View File
@@ -92,6 +92,8 @@ final: default
There may be additional keys depending on the filter type.
If you have [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) available, you can use any blend mode it supports. Otherwise OCS provides these built-in blend modes: `lerp`, `a_only`, `b_only`. _Note_: `a` is considered the original value, `b` the changed value. `a_only` and `b_only` will still scale their output by the `strength`.
## Filter Types
### `simple`
+9 -8
View File
@@ -2,17 +2,18 @@ from . import types, expression, handler, util, validation
from .expression import Expression
from .validation import Arg, ValidateArg
from .handler import BASIC_HANDLERS, BaseHandler
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
__all__ = (
"types",
"expression",
"handler",
"util",
"validation",
"ValidateArg",
"Expression",
"Arg",
"BaseHandler",
"BASIC_HANDLERS",
"expression",
"Expression",
"handler",
"HandlerContext",
"types",
"util",
"ValidateArg",
"validation",
)
+46 -20
View File
@@ -15,13 +15,15 @@ from .types import (
ExpKV,
)
COMMA_PRECEDENCE = 2
class Expression:
EXPR_RE = re.compile(
r"""
\s*
(
\d+ # Possibly negative numeric literal
\d+ # Numeric literal
(?: \. \d* )? # Floating point
(?: e [+-] \d+)? # Scientific notation
| (?: \*\* | // ) # Doubled operators
@@ -30,9 +32,11 @@ class Expression:
| (?: \|\| | && ) # Logic
| [-+*/|!(),] # Operators
| :> # Key value binop
| ;
| \[ | ]
| \.\.\.
| := # Assignment
| ; # Sequencing
| [?:] # Ternary
| \[ | ] # Index
| \.\.\. # Index ellipsis
| '[\w.]+ # Symbol
| `?[a-z][\w.]*`? # Function/variable names
)
@@ -53,7 +57,8 @@ class Expression:
return self.eval(*args, **kwargs)
def eval(self, handlers, *args, **kwargs):
print("\nEVAL", self.expr)
if self.expr != ExpOp("default"):
print("\nEVAL", self.expr)
if not isinstance(self.expr, ExpBase):
return self.expr
return self.expr.eval(handlers, *args, **kwargs)
@@ -93,7 +98,7 @@ class Expression:
yield from (cls.fixup_token(m.group(1)) for m in cls.EXPR_RE.finditer(s))
STATIC_OP_HANDLERS = {
CONST_OP_HANDLERS = {
"+": operator.add,
"-": operator.sub,
"*": operator.mul,
@@ -108,18 +113,29 @@ STATIC_OP_HANDLERS = {
"idiv": operator.floordiv,
"pow": operator.pow,
"mod": operator.mod,
"neg": operator.neg,
">": operator.gt,
"<": operator.lt,
">=": operator.ge,
"<=": operator.le,
"!=": operator.ne,
"==": operator.eq,
}
def is_const_value(val):
return val in (None, True, False) or isinstance(val, (int, float, ExpSym))
def make_funap(op, args=(), kwargs=None):
if kwargs is None:
kwargs = ExpDict()
argc = len(args)
if argc > 2 or len(kwargs) or not all(isinstance(v, (int, float)) for v in args):
if argc > 2 or len(kwargs) or not all(is_const_value(v) for v in args):
return ExpFunAp(op, args, kwargs)
if argc == 1 and op in "-+":
return -args[0] if op == "-" else args[0]
h = STATIC_OP_HANDLERS.get(op)
h = CONST_OP_HANDLERS.get(op)
if h is None:
return ExpFunAp(op, args, kwargs)
return h(*args)
@@ -171,7 +187,7 @@ class ExprParserSpec(ParserSpec):
raise ParseError(f"{left!r} is not a valid function/variable name")
args = []
while p.lexer and p.token != ")":
args.append(p.parse_until(1))
args.append(p.parse_until(COMMA_PRECEDENCE))
if p.token == ",":
p.advance()
p.expect(")")
@@ -186,15 +202,9 @@ class ExprParserSpec(ParserSpec):
@staticmethod
def left_semicolon(p, token, left, bp):
if p.token == ")" or p.token is None:
return (
left
if isinstance(left, ExpStatements)
else ExpStatements(ExpTuple((left,)))
)
r = p.parse_until(bp)
r = None if p.token in (None, ")", ";") else p.parse_until(0)
return ExpStatements(
ExpTuple(*left.statements, r)
ExpTuple((*left.statements, r))
if isinstance(left, ExpStatements)
else ExpTuple((left, r))
)
@@ -205,6 +215,20 @@ class ExprParserSpec(ParserSpec):
p.expect("]")
return make_funap("index", ExpTuple((idx, left)))
@staticmethod
def left_assign(p, token, left, bp):
if not isinstance(left, (ExpOp, ExpSym)):
raise ParseError(f"bad LHS type for assignment operation {type(left)}")
val = p.parse_until(bp)
return make_funap("set_var", ExpTuple((ExpSym(left), val)))
@staticmethod
def left_ternary(p, token, left, bp):
true_branch = p.parse_until(0)
p.expect(":")
false_branch = p.parse_until(bp)
return make_funap("if", ExpTuple((left, true_branch, false_branch)))
@staticmethod
def get_type(token):
if isinstance(token, (int, float)):
@@ -230,10 +254,12 @@ class ExprParserSpec(ParserSpec):
self.add_left(9, self.left_binop, ("&&",))
self.add_left(7, self.left_binop, ("||",))
self.add_left(6, self.left_kv, (":>",))
self.add_left(5, self.left_semicolon, (";",))
self.add_left(1, self.left_comma, (",",))
self.add_leftright(5, self.left_ternary, ("?",))
self.add_leftright(4, self.left_assign, (":=",))
self.add_left(COMMA_PRECEDENCE, self.left_comma, (",",))
self.add_left(1, self.left_semicolon, (";",))
self.add_null(0, self.null_paren, ("(",))
self.add_null(
-1, self.null_constant, ("number", "op", "sym", Ellipsis, True, False, None)
)
self.add_null(-1, ParserSpec.null_error, (")", "]"))
self.add_null(-1, ParserSpec.null_error, (")", "]", ":"))
+58 -7
View File
@@ -1,7 +1,7 @@
import operator
from .validation import ValidateArg, Arg, ValidateError
from .types import Empty, ExpDict
from .types import Empty, ExpDict, ExpOp
from .util import torch
@@ -9,6 +9,47 @@ class HandlerError(Exception):
pass
class HandlerContext:
def __init__(self, handlers=None, constants=None, variables=None):
self.handlers = handlers if handlers is not None else {}
self.constants = constants if constants is not None else {}
self.variables = variables if variables is not None else {}
def get_handler(self, k, default=Empty):
return self.handlers.get(k, default)
def get_var(self, k, default=Empty):
result = self.constants.get(k, Empty)
if result is Empty:
result = self.variables.get(k, Empty)
return default if result is Empty else result
def set_var(self, k, v):
if k in self.constants:
raise KeyError(
f"Cannot set variable with key {k}: already exists as a constant"
)
self.variables[k] = v
def unset_var(self, k):
if k in self.variables:
del self.variables[k]
return True
return False
def __contains__(self, k):
return any(
k in coll for coll in (self.handlers, self.constants, self.variables)
)
def clone(self, *, handlers=Empty, constants=Empty, variables=Empty):
return self.__class__(
self.handlers if handlers is Empty else handlers,
self.constants if constants is Empty else constants,
self.variables if variables is Empty else variables,
)
class BaseHandler:
input_validators = ()
@@ -50,7 +91,7 @@ class BaseHandler:
str_eff_key = True
else:
raise ValidateError(
f"Error validating input argument {key}, out of range for actual function arguments"
f"Error validating input argument {key} for {obj.name}, out of range for actual function arguments"
)
if getter is None:
if str_eff_key:
@@ -65,7 +106,7 @@ class BaseHandler:
return validator(key, val)
except ValidateError as exc:
raise ValidateError(
f"Error validating input argument {key}, type {type(val)}: {exc!r}"
f"Error validating input argument {key} for {obj.name}, type {type(val)}: {exc!r}"
) from None
def safe_get_multi(self, keys, obj, getter=None, *, default=Empty):
@@ -218,7 +259,7 @@ class IsSetHandler(BaseHandler):
def handle(self, obj, getter):
key = self.safe_get(0, obj, getter=getter)
return key in getter.handlers
return key in getter.ctx
def validate_output(self, obj, value):
return operator.truth(value)
@@ -232,10 +273,10 @@ class GetHandler(BaseHandler):
def handle(self, obj, getter):
key = self.safe_get("name", obj, getter=getter)
h = getter.handlers.get(key)
if h is None:
result = getter.ctx.get_var(key)
if result is Empty:
return self.safe_get("fallback", obj, getter=getter)
return h(getter.handlers, *getter.args, **getter.kwargs)
return ExpOp(key).eval(getter.ctx, *getter.args, **getter.kwargs)
class S_Handler(BaseHandler):
@@ -305,6 +346,15 @@ class CommentHandler(BaseHandler):
return None
class SetVarHandler(BaseHandler):
input_validators = (Arg.string("lhs"), Arg.present("rhs"))
def handle(self, obj, getter):
key, val = self.safe_get_all(obj, getter)
getter.ctx.set_var(key, val)
return val
LOGIC_HANDLERS = {
"||": OrHandler(),
"&&": AndHandler(),
@@ -359,6 +409,7 @@ MISC_HANDLERS = {
"unsafe_call": UnsafeCallHandler(),
"dict": DictHandler(),
"comment": CommentHandler(),
"set_var": SetVarHandler(),
}
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
+11 -18
View File
@@ -21,10 +21,10 @@ class ExpOp(str, ExpBase):
__slots__ = ()
def eval(self, handlers, *args, **kwargs):
h = handlers.get(self)
if h is None:
value = handlers.get_var(self)
if value is Empty:
raise KeyError(f"No handler for op/var {self}")
return h(handlers, *args, **kwargs)
return value
class ExpBinOp(ExpOp):
@@ -135,27 +135,20 @@ class ExpStatements(ExpBase):
class ExprGetter:
GetterEmpty = Empty
# class GetterEmpty:
# def __bool__(self):
# return False
def __init__(self, obj, handlers, *args, **kwargs):
def __init__(self, obj, ctx, *args, **kwargs):
self.obj = obj
self.handlers = handlers
self.ctx = ctx
self.args = args
self.kwargs = kwargs
def __call__(self, k, *, default=GetterEmpty):
def __call__(self, k, *, default=Empty):
obj = self.obj
result = (
obj.kwargs.get_eval(
k, self.handlers, *self.args, default=default, **self.kwargs
)
obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs)
if isinstance(k, str)
else obj.args.get_eval(k, self.handlers, *self.args, **self.kwargs)
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
)
if result is self.GetterEmpty:
if result is Empty:
raise KeyError(f"Unknown key {k!r}")
return result
@@ -169,8 +162,8 @@ class ExpFunAp(ExpBase):
self.kwargs = kwargs if kwargs is not None else ExpDict()
def eval(self, handlers, *args, **kwargs):
handler = handlers.get(self.name)
if handler is None:
handler = handlers.get_handler(self.name)
if handler is Empty:
raise KeyError(f"No handler for op: {self.name!r}")
return handler(
self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs
+3 -6
View File
@@ -1,14 +1,12 @@
import functools
from .util import torch
from .types import Empty
class Arg:
__slots__ = ("name", "default", "validator")
class Empty:
pass
def __init__(self, name, default=Empty, *, validator=None):
self.name = name
self.default = default
@@ -18,9 +16,8 @@ class Arg:
return self.validate(value, *args, **kwargs)
def validate(self, value):
# FIXME: This shouldn't be using None.
if value is None:
if self.default is self.Empty:
if value is Empty:
if self.default is Empty:
raise ValueError(f"Missing value for argument {self.name}")
return self.default
try:
+1 -1
View File
@@ -173,7 +173,7 @@ class NoiseHandler(NormHandler):
def handle(self, obj, getter):
t, typ = self.safe_get_all(obj, getter)
ctx = getter.handlers
ctx = getter.ctx
smin, smax, s, sn = (
h(ctx, *getter.args, **getter.kwargs) if h is not None else None
for h in (
+4 -31
View File
@@ -28,35 +28,8 @@ BLENDING_MODES = BLENDING_MODES | {
FILTER = {}
class FilterHandlerCollection:
def __init__(self, op_handlers, refs):
self.op_handlers = op_handlers
self.refs = refs
def get(self, k, default=None):
result = self.op_handlers.get(k)
if result is not None:
return result
result = self.refs.get(k)
if result is None:
return None
def h(obj, *_args, __k=k, __val=result, **_kwargs):
if not isinstance(obj, FilterHandlerCollection):
raise ValueError(f"Unexpected arguments to variable reference {k}")
return result
return h
def __contains__(self, k):
return k in self.op_handlers or k in self.refs
def clone_with_refs(self, refs):
return self.__class__(self.op_handlers, refs)
FILTER_HANDLERS = FilterHandlerCollection(
expr.BASIC_HANDLERS | expression_handlers.HANDLERS, {}
FILTER_HANDLERS = expr.HandlerContext(
expr.BASIC_HANDLERS | expression_handlers.HANDLERS
)
@@ -240,7 +213,7 @@ class Filter:
if self.when is None:
return True
refs = fallback(refs, FilterRefs())
matched = self.when.eval(FILTER_HANDLERS.clone_with_refs(refs))
matched = self.when.eval(FILTER_HANDLERS.clone(constants=refs, variables={}))
# if matched:
# print("\nMATCH", self.name)
return matched
@@ -250,7 +223,7 @@ class Filter:
return ops.apply(default_ref, ops, refs=refs)
drefs = FilterRefs({"default": default_ref})
refs = drefs if refs is None else refs | drefs
return ops.eval(FILTER_HANDLERS.clone_with_refs(refs))
return ops.eval(FILTER_HANDLERS.clone(constants=refs, variables={}))
class SimpleFilter(Filter):
+2 -1
View File
@@ -14,7 +14,8 @@ def find_merge_sampler(merge_samplers, ss) -> object | None:
handlers = None
for merge_sampler in merge_samplers:
if merge_sampler.when is not None and handlers is None:
handlers = FILTER_HANDLERS.clone_with_refs(ss.refs)
handlers = FILTER_HANDLERS.clone(constants=ss.refs)
# handlers = FILTER_HANDLERS.clone_with_refs(ss.refs)
if merge_sampler.check_match(handlers, ss=ss):
return merge_sampler
return None
+5 -3
View File
@@ -382,10 +382,12 @@ class EulerStep(SingleStepSampler):
class CycleSingleStepSampler(SingleStepSampler):
def __init__(self, *, cycle_pct=0.25, **kwargs):
super().__init__(**kwargs)
if cycle_pct < 0:
raise ValueError("cycle_pct must be positive")
self.cycle_pct = cycle_pct
def get_cycle_scales(self, sigma_next):
keep_scale = sigma_next * (1.0 - self.cycle_pct)
keep_scale = sigma_next * (1.0 - self.cycle_pct) if self.cycle_pct < 1 else 0.0
add_scale = ((sigma_next**2.0 - keep_scale**2.0) ** 0.5) * (
0.95 + 0.25 * self.cycle_pct
)
@@ -401,9 +403,9 @@ class EulerCycleStep(CycleSingleStepSampler):
def step(self, x, ss):
if ss.sigma_next == 0:
return (yield from self.denoised_result(ss))
d = self.to_d(ss.hcur)
keep_scale, add_scale = self.get_cycle_scales(ss.sigma_next)
yield from self.result(ss, ss.denoised + d * keep_scale, add_scale)
keep_noise = self.to_d(ss.hcur) * keep_scale if keep_scale > 0 else 0.0
yield from self.result(ss, ss.denoised + keep_noise, add_scale)
class DPMPP2MStep(HistorySingleStepSampler, DPMPPStepMixin):