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):