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:
@@ -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`
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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
@@ -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, (")", "]", ":"))
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user