From 2481e7263b5779ef297ae0cfdaa7b3b876df7de6 Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 8 Aug 2024 06:27:53 -0600 Subject: [PATCH] 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 --- README.md | 9 +++++ docs/expression.md | 6 ++++ docs/filter.md | 2 ++ py/expression/__init__.py | 17 +++++----- py/expression/expression.py | 66 ++++++++++++++++++++++++++----------- py/expression/handler.py | 65 ++++++++++++++++++++++++++++++++---- py/expression/types.py | 29 +++++++--------- py/expression/validation.py | 9 ++--- py/expression_handlers.py | 2 +- py/filtering.py | 35 +++----------------- py/sampling.py | 3 +- py/step_samplers.py | 8 +++-- 12 files changed, 156 insertions(+), 95 deletions(-) diff --git a/README.md b/README.md index 4bbea1f..315acab 100644 --- a/README.md +++ b/README.md @@ -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` diff --git a/docs/expression.md b/docs/expression.md index 922ed82..5f8482c 100644 --- a/docs/expression.md +++ b/docs/expression.md @@ -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. | Boolean negation | |⬤| `s_` | start:`I(null)`, end:`I(null)`, step:`I(null)` | `slice` | | Creates a slice object from the `start`, `end`, `step` values. See Numpy [s_](https://numpy.org/doc/stable/reference/generatednumpy.s_.html) | + |⬤| `set_var` | `SY`, `*` | `*` | + | Sets a temporary variable to the specified value and returns the value. Alias for the `:=` assignment operator.
**Example**: `test1 := 2; set_var('test2, 10); test1 * test2` | |⬤| `unsafe_call` | `callable`, `*`\* | `*` | | Allows calling an arbitrary callable.
**Example:** `unsafe_call(some_callable, arg1, arg2, kwarg1 :> 123)` diff --git a/docs/filter.md b/docs/filter.md index c7a0e9f..21197b9 100644 --- a/docs/filter.md +++ b/docs/filter.md @@ -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` diff --git a/py/expression/__init__.py b/py/expression/__init__.py index 07fc6c9..2ef5a04 100644 --- a/py/expression/__init__.py +++ b/py/expression/__init__.py @@ -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", ) diff --git a/py/expression/expression.py b/py/expression/expression.py index df3c836..3a09578 100644 --- a/py/expression/expression.py +++ b/py/expression/expression.py @@ -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, (")", "]", ":")) diff --git a/py/expression/handler.py b/py/expression/handler.py index 2a5d83a..a647d12 100644 --- a/py/expression/handler.py +++ b/py/expression/handler.py @@ -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 diff --git a/py/expression/types.py b/py/expression/types.py index 8172452..71f048e 100644 --- a/py/expression/types.py +++ b/py/expression/types.py @@ -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 diff --git a/py/expression/validation.py b/py/expression/validation.py index b7ff72c..d01ca39 100644 --- a/py/expression/validation.py +++ b/py/expression/validation.py @@ -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: diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 3f01b3f..fb0d6af 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -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 ( diff --git a/py/filtering.py b/py/filtering.py index 77a8c81..6ef2a71 100644 --- a/py/filtering.py +++ b/py/filtering.py @@ -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): diff --git a/py/sampling.py b/py/sampling.py index 48b44bc..0f419cb 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -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 diff --git a/py/step_samplers.py b/py/step_samplers.py index c43c079..fe409a4 100644 --- a/py/step_samplers.py +++ b/py/step_samplers.py @@ -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):