Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
79c7b24d1e | ||
|
|
5ea059bed5 | ||
|
|
ee59df94e3 |
@@ -11,5 +11,7 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
|
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
|
||||||
"OCS ApplyFilterLatent": nodes.ApplyFilterLatent,
|
"OCS ApplyFilterLatent": nodes.ApplyFilterLatent,
|
||||||
"OCS ApplyFilterImage": nodes.ApplyFilterImage,
|
"OCS ApplyFilterImage": nodes.ApplyFilterImage,
|
||||||
|
"OCS ExpressionFilteredLatentOperation": nodes.ExpressionFilteredLatentOperationNode,
|
||||||
|
"OCS ExpressionFilteredModelPatch": nodes.ExpressionFilteredModelPatchNode,
|
||||||
} | custom_noise.NODE_CLASS_MAPPINGS
|
} | custom_noise.NODE_CLASS_MAPPINGS
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||||
|
|||||||
@@ -1,11 +1,11 @@
|
|||||||
import abc
|
import abc
|
||||||
|
from typing import Any, Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from typing import Callable, Any
|
|
||||||
|
|
||||||
from ..external import IntegratedNode
|
from ..external import IntegratedNode
|
||||||
|
from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE
|
||||||
from ..noise import scale_noise
|
from ..noise import scale_noise
|
||||||
from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT
|
|
||||||
|
|
||||||
|
|
||||||
class CustomNoiseItemBase(abc.ABC):
|
class CustomNoiseItemBase(abc.ABC):
|
||||||
|
|||||||
+605
-244
File diff suppressed because it is too large
Load Diff
+578
-419
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,12 @@
|
|||||||
from . import types, expression, handler, util, validation
|
from . import expression, types, util
|
||||||
|
|
||||||
from .expression import Expression
|
from .expression import Expression
|
||||||
from .validation import Arg, ValidateArg
|
|
||||||
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
|
try:
|
||||||
|
from . import handler, validation
|
||||||
|
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
|
||||||
|
from .validation import Arg, ValidateArg
|
||||||
|
except (ImportError, ModuleNotFoundError):
|
||||||
|
pass
|
||||||
|
|
||||||
__all__ = (
|
__all__ = (
|
||||||
"Arg",
|
"Arg",
|
||||||
|
|||||||
+27
-12
@@ -1,20 +1,22 @@
|
|||||||
import re
|
|
||||||
import operator
|
import operator
|
||||||
|
import re
|
||||||
|
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from .parser import Parser, ParserSpec, ParseError
|
from .parser import ParseError, Parser, ParserSpec
|
||||||
from .types import (
|
from .types import (
|
||||||
Empty,
|
Empty,
|
||||||
ExpBase,
|
ExpBase,
|
||||||
ExpOp,
|
|
||||||
ExpBinOp,
|
ExpBinOp,
|
||||||
ExpSym,
|
|
||||||
ExpStatements,
|
|
||||||
ExpFunAp,
|
|
||||||
ExpTuple,
|
|
||||||
ExpDict,
|
ExpDict,
|
||||||
|
ExpFunAp,
|
||||||
ExpKV,
|
ExpKV,
|
||||||
|
ExpMethodAp,
|
||||||
|
ExpOp,
|
||||||
|
ExpReturn,
|
||||||
|
ExpStatements,
|
||||||
|
ExpSym,
|
||||||
|
ExpTuple,
|
||||||
)
|
)
|
||||||
|
|
||||||
COMMA_PRECEDENCE = 2
|
COMMA_PRECEDENCE = 2
|
||||||
@@ -36,10 +38,11 @@ class Expression:
|
|||||||
| :> # Key value binop
|
| :> # Key value binop
|
||||||
| := # Assignment
|
| := # Assignment
|
||||||
| ; # Sequencing
|
| ; # Sequencing
|
||||||
|
| :: # Method call
|
||||||
| [?:] # Ternary
|
| [?:] # Ternary
|
||||||
| \[ | ] # Index
|
| \[ | ] # Index
|
||||||
| \.\.\. # Index ellipsis
|
| \.\.\. # Index ellipsis
|
||||||
| '[-\w.]+ # Symbol
|
| '[-\w.:=]+ # Symbol
|
||||||
| `?[a-z][\w.]*`? # Function/variable names
|
| `?[a-z][\w.]*`? # Function/variable names
|
||||||
)
|
)
|
||||||
\s*
|
\s*
|
||||||
@@ -63,7 +66,10 @@ class Expression:
|
|||||||
tqdm.write(f"* OCS: EVAL: {self.expr}")
|
tqdm.write(f"* OCS: EVAL: {self.expr}")
|
||||||
if not isinstance(self.expr, ExpBase):
|
if not isinstance(self.expr, ExpBase):
|
||||||
return self.expr
|
return self.expr
|
||||||
return self.expr.eval(handlers, *args, **kwargs)
|
try:
|
||||||
|
return self.expr.eval(handlers, *args, **kwargs)
|
||||||
|
except ExpReturn as expret:
|
||||||
|
return expret.args[0]
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
return len(self.expr)
|
return len(self.expr)
|
||||||
@@ -157,9 +163,9 @@ class ExprParserSpec(ParserSpec):
|
|||||||
def split_funap_args(toks):
|
def split_funap_args(toks):
|
||||||
if not isinstance(toks, (list, tuple)):
|
if not isinstance(toks, (list, tuple)):
|
||||||
return ExpTuple((toks,)), ExpDict()
|
return ExpTuple((toks,)), ExpDict()
|
||||||
return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict({
|
return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict(
|
||||||
str(t.k): t.v for t in toks if isinstance(t, ExpKV)
|
{str(t.k): t.v for t in toks if isinstance(t, ExpKV)}
|
||||||
})
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def null_constant(p, token, bp):
|
def null_constant(p, token, bp):
|
||||||
@@ -200,6 +206,14 @@ class ExprParserSpec(ParserSpec):
|
|||||||
p.expect(")")
|
p.expect(")")
|
||||||
return make_funap(left, *cls.split_funap_args(args))
|
return make_funap(left, *cls.split_funap_args(args))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def left_methodcall(cls, p, token, left, bp):
|
||||||
|
methname = p.parse_until(31)
|
||||||
|
p.expect("(")
|
||||||
|
funap = cls.left_funcall(p, token=None, left=methname, bp=None)
|
||||||
|
funap.args = ExpTuple((Empty, *funap.args))
|
||||||
|
return ExpMethodAp(left, funap)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def left_comma(p, token, left, bp):
|
def left_comma(p, token, left, bp):
|
||||||
if p.token == ")":
|
if p.token == ")":
|
||||||
@@ -251,6 +265,7 @@ class ExprParserSpec(ParserSpec):
|
|||||||
def populate(self):
|
def populate(self):
|
||||||
self.add_left(31, self.left_funcall, ("(",))
|
self.add_left(31, self.left_funcall, ("(",))
|
||||||
self.add_left(31, self.left_index, ("[",))
|
self.add_left(31, self.left_index, ("[",))
|
||||||
|
self.add_left(31, self.left_methodcall, ("::",))
|
||||||
self.add_leftright(29, self.left_binop, ("**",))
|
self.add_leftright(29, self.left_binop, ("**",))
|
||||||
self.add_null(27, self.null_prefixop, ("+", "-", "!"))
|
self.add_null(27, self.null_prefixop, ("+", "-", "!"))
|
||||||
self.add_left(25, self.left_binop, ("*", "/"))
|
self.add_left(25, self.left_binop, ("*", "/"))
|
||||||
|
|||||||
+105
-15
@@ -1,8 +1,11 @@
|
|||||||
import operator
|
import operator
|
||||||
|
import traceback
|
||||||
|
|
||||||
from .validation import ValidateArg, Arg, ValidateError
|
from tqdm import tqdm
|
||||||
from .types import Empty, ExpDict, ExpOp
|
|
||||||
|
from .types import Empty, ExpDict, ExpOp, ExpReturn, ExpTuple
|
||||||
from .util import torch
|
from .util import torch
|
||||||
|
from .validation import Arg, ValidateArg, ValidateError
|
||||||
|
|
||||||
|
|
||||||
class HandlerError(Exception):
|
class HandlerError(Exception):
|
||||||
@@ -61,9 +64,12 @@ class BaseHandler:
|
|||||||
def __call__(self, obj, *, getter):
|
def __call__(self, obj, *, getter):
|
||||||
try:
|
try:
|
||||||
val = self.handle(obj, getter)
|
val = self.handle(obj, getter)
|
||||||
return self.validate_output(obj, val)
|
except ExpReturn:
|
||||||
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
raise HandlerError(f'Error evaluating "{obj.name}":\n {exc!r}') from exc
|
tb = traceback.format_exc()
|
||||||
|
raise HandlerError(f'Error evaluating "{obj.name}": {exc!s}\n{tb}') from exc
|
||||||
|
return self.validate_output(obj, val)
|
||||||
|
|
||||||
def safe_get(self, key, obj, getter=None, *, default=Empty):
|
def safe_get(self, key, obj, getter=None, *, default=Empty):
|
||||||
str_key = isinstance(key, str)
|
str_key = isinstance(key, str)
|
||||||
@@ -256,6 +262,22 @@ class UnarySimpleMathHandler(SimpleMathHandler):
|
|||||||
input_validators = (Arg.numeric("lhs"),)
|
input_validators = (Arg.numeric("lhs"),)
|
||||||
|
|
||||||
|
|
||||||
|
class SimpleOpHandler(BaseHandler):
|
||||||
|
input_validators = (Arg.present("lhs"), Arg.present("rhs"))
|
||||||
|
|
||||||
|
def __init__(self, handler):
|
||||||
|
super().__init__()
|
||||||
|
self.handler = handler
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
args = (
|
||||||
|
self.safe_get(idx, obj, getter=getter)
|
||||||
|
for idx in range(len(self.input_validators))
|
||||||
|
)
|
||||||
|
result = self.handler(*args)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
class IsSetHandler(BaseHandler):
|
class IsSetHandler(BaseHandler):
|
||||||
input_validators = (Arg.string("name"),)
|
input_validators = (Arg.string("name"),)
|
||||||
|
|
||||||
@@ -328,6 +350,13 @@ class MaxHandler(MinHandler):
|
|||||||
return max(*self.safe_get("values", obj, getter))
|
return max(*self.safe_get("values", obj, getter))
|
||||||
|
|
||||||
|
|
||||||
|
class SumHandler(BaseHandler):
|
||||||
|
input_validators = (Arg.numeric_sequence("values"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
return sum(tuple(self.safe_get("values", obj, getter)))
|
||||||
|
|
||||||
|
|
||||||
class UnsafeCallHandler(BaseHandler):
|
class UnsafeCallHandler(BaseHandler):
|
||||||
input_validators = (Arg.present("__callable"),)
|
input_validators = (Arg.present("__callable"),)
|
||||||
|
|
||||||
@@ -338,7 +367,7 @@ class UnsafeCallHandler(BaseHandler):
|
|||||||
)
|
)
|
||||||
fun = self.safe_get("__callable", obj, getter)
|
fun = self.safe_get("__callable", obj, getter)
|
||||||
if not callable(fun):
|
if not callable(fun):
|
||||||
raise ValueError("Cannot call supplied value: not a callable")
|
raise TypeError("Cannot call supplied value: not a callable")
|
||||||
args = (self.safe_get(idx, obj, getter) for idx in range(1, len(obj.args)))
|
args = (self.safe_get(idx, obj, getter) for idx in range(1, len(obj.args)))
|
||||||
kwargs = {k: self.safe_get(k, obj, getter) for k in obj.kwargs}
|
kwargs = {k: self.safe_get(k, obj, getter) for k in obj.kwargs}
|
||||||
return fun(*args, **kwargs)
|
return fun(*args, **kwargs)
|
||||||
@@ -365,6 +394,53 @@ class SetVarHandler(BaseHandler):
|
|||||||
return val
|
return val
|
||||||
|
|
||||||
|
|
||||||
|
class ReturnHandler(BaseHandler):
|
||||||
|
input_validators = (Arg.present("expression"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
raise ExpReturn(self.safe_get("expression", obj, getter))
|
||||||
|
|
||||||
|
|
||||||
|
class PrintHandler(BaseHandler):
|
||||||
|
input_validators = (Arg.present("lhs"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
lhs = self.safe_get("lhs", obj, getter)
|
||||||
|
tqdm.write(f"[OCS expr_print]: {lhs!s}")
|
||||||
|
|
||||||
|
|
||||||
|
class MapHandler(BaseHandler):
|
||||||
|
input_validators = (
|
||||||
|
Arg.sequence("items"),
|
||||||
|
Arg.string("key", default="item"),
|
||||||
|
Arg.present("expression"),
|
||||||
|
Arg.present("check_expression", default=None),
|
||||||
|
)
|
||||||
|
|
||||||
|
class _MapEmpty:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
items = self.safe_get("items", obj, getter)
|
||||||
|
key = self.safe_get("key", obj, getter)
|
||||||
|
have_check_expr = None
|
||||||
|
result = []
|
||||||
|
for item in items:
|
||||||
|
getter.ctx.set_var(key, item)
|
||||||
|
if have_check_expr in {True, None}:
|
||||||
|
checked = self.safe_get(
|
||||||
|
"check_expression",
|
||||||
|
obj,
|
||||||
|
getter,
|
||||||
|
default=self._MapEmpty,
|
||||||
|
)
|
||||||
|
have_check_expr = checked is not self._MapEmpty
|
||||||
|
if have_check_expr and not bool(checked):
|
||||||
|
continue
|
||||||
|
result.append(self.safe_get("expression", obj, getter))
|
||||||
|
return ExpTuple(result)
|
||||||
|
|
||||||
|
|
||||||
LOGIC_HANDLERS = {
|
LOGIC_HANDLERS = {
|
||||||
"||": OrHandler(),
|
"||": OrHandler(),
|
||||||
"&&": AndHandler(),
|
"&&": AndHandler(),
|
||||||
@@ -385,21 +461,26 @@ for k, alias in (
|
|||||||
|
|
||||||
|
|
||||||
MATH_HANDLERS = {
|
MATH_HANDLERS = {
|
||||||
|
"*": SimpleMathHandler(operator.mul),
|
||||||
|
"**": SimpleMathHandler(operator.pow),
|
||||||
"+": SimpleMathHandler(operator.add),
|
"+": SimpleMathHandler(operator.add),
|
||||||
"-": MinusHandler(),
|
"-": MinusHandler(),
|
||||||
"*": SimpleMathHandler(operator.mul),
|
|
||||||
"/": SimpleMathHandler(operator.truediv),
|
"/": SimpleMathHandler(operator.truediv),
|
||||||
"//": SimpleMathHandler(operator.floordiv),
|
"//": SimpleMathHandler(operator.floordiv),
|
||||||
"**": SimpleMathHandler(operator.pow),
|
|
||||||
"mod": SimpleMathHandler(operator.mod),
|
|
||||||
"neg": UnarySimpleMathHandler(operator.neg),
|
|
||||||
"between": BetweenHandler(),
|
|
||||||
"<": RelComparisonHandler(operator.lt),
|
"<": RelComparisonHandler(operator.lt),
|
||||||
"<=": RelComparisonHandler(operator.le),
|
"<=": RelComparisonHandler(operator.le),
|
||||||
">": RelComparisonHandler(operator.gt),
|
">": RelComparisonHandler(operator.gt),
|
||||||
">=": RelComparisonHandler(operator.ge),
|
">=": RelComparisonHandler(operator.ge),
|
||||||
"min": MinHandler(),
|
"abs": SimpleMathHandler(operator.abs),
|
||||||
|
"between": BetweenHandler(),
|
||||||
|
"bool": UnarySimpleMathHandler(handler=bool),
|
||||||
|
"float": UnarySimpleMathHandler(handler=float),
|
||||||
|
"int": UnarySimpleMathHandler(handler=int),
|
||||||
"max": MaxHandler(),
|
"max": MaxHandler(),
|
||||||
|
"min": MinHandler(),
|
||||||
|
"mod": SimpleMathHandler(operator.mod),
|
||||||
|
"neg": UnarySimpleMathHandler(operator.neg),
|
||||||
|
"sum": SumHandler(),
|
||||||
}
|
}
|
||||||
for k, alias in (
|
for k, alias in (
|
||||||
("+", "add"),
|
("+", "add"),
|
||||||
@@ -412,14 +493,23 @@ for k, alias in (
|
|||||||
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
|
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
|
||||||
|
|
||||||
MISC_HANDLERS = {
|
MISC_HANDLERS = {
|
||||||
"is_set": IsSetHandler(),
|
"and": SimpleOpHandler(operator.and_),
|
||||||
|
"comment": CommentHandler(),
|
||||||
|
"concat": SimpleOpHandler(operator.concat),
|
||||||
|
"contains": SimpleOpHandler(operator.contains),
|
||||||
|
"dict": DictHandler(),
|
||||||
"get": GetHandler(),
|
"get": GetHandler(),
|
||||||
"index": IndexHandler(),
|
"index": IndexHandler(),
|
||||||
|
"is_set": IsSetHandler(),
|
||||||
|
"map": MapHandler(),
|
||||||
|
"op_or": SimpleOpHandler(operator.or_),
|
||||||
|
"op_and": SimpleOpHandler(operator.and_),
|
||||||
|
"op_xor": SimpleOpHandler(operator.xor),
|
||||||
|
"print": PrintHandler(),
|
||||||
|
"return": ReturnHandler(),
|
||||||
"s_": S_Handler(),
|
"s_": S_Handler(),
|
||||||
"unsafe_call": UnsafeCallHandler(),
|
|
||||||
"dict": DictHandler(),
|
|
||||||
"comment": CommentHandler(),
|
|
||||||
"set_var": SetVarHandler(),
|
"set_var": SetVarHandler(),
|
||||||
|
"unsafe_call": UnsafeCallHandler(),
|
||||||
}
|
}
|
||||||
|
|
||||||
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|
||||||
|
|||||||
+98
-30
@@ -3,6 +3,10 @@ class Empty:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class ExpReturn(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class ExpBase:
|
class ExpBase:
|
||||||
def __bool__(self):
|
def __bool__(self):
|
||||||
return True
|
return True
|
||||||
@@ -41,8 +45,10 @@ class ExpSym(str, ExpBase):
|
|||||||
class ExpTuple(tuple, ExpBase):
|
class ExpTuple(tuple, ExpBase):
|
||||||
__slots__ = ()
|
__slots__ = ()
|
||||||
|
|
||||||
def clone(self):
|
def clone(self, **kwargs):
|
||||||
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
|
return self.__class__(
|
||||||
|
v.clone(**kwargs) if isinstance(v, ExpBase) else v for v in self
|
||||||
|
)
|
||||||
|
|
||||||
def get_eval(self, k, handlers, *args, default=None, **kwargs):
|
def get_eval(self, k, handlers, *args, default=None, **kwargs):
|
||||||
val = super().__getitem__(k)
|
val = super().__getitem__(k)
|
||||||
@@ -77,11 +83,10 @@ class ExpKV(ExpBase):
|
|||||||
class ExpDict(dict, ExpBase):
|
class ExpDict(dict, ExpBase):
|
||||||
__slots__ = ()
|
__slots__ = ()
|
||||||
|
|
||||||
def clone(self):
|
def clone(self, **kwargs):
|
||||||
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
|
return self.__class__(
|
||||||
|
v.clone(**kwargs) if isinstance(v, ExpBase) else v for v in self
|
||||||
def pop(self, *args, **kwargs):
|
)
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def get_eval(self, k, handlers, *args, default=Empty, **kwargs):
|
def get_eval(self, k, handlers, *args, default=Empty, **kwargs):
|
||||||
val = super().get(k, default)
|
val = super().get(k, default)
|
||||||
@@ -106,12 +111,17 @@ class ExpDict(dict, ExpBase):
|
|||||||
for k, v in self.items()
|
for k, v in self.items()
|
||||||
}
|
}
|
||||||
|
|
||||||
popitem = pop
|
# Can't remember if there was a compelling reason ExpDict can't be mutable but
|
||||||
update = pop
|
# it breaks deep copy stuff.
|
||||||
clear = pop
|
#
|
||||||
__delitem__ = pop
|
# def pop(self, *args, **kwargs):
|
||||||
__setitem__ = pop
|
# raise NotImplementedError
|
||||||
__ior__ = pop
|
# popitem = pop
|
||||||
|
# update = pop
|
||||||
|
# clear = pop
|
||||||
|
# __delitem__ = pop
|
||||||
|
# __setitem__ = pop
|
||||||
|
# __ior__ = pop
|
||||||
|
|
||||||
|
|
||||||
class ExpStatements(ExpBase):
|
class ExpStatements(ExpBase):
|
||||||
@@ -135,26 +145,81 @@ class ExpStatements(ExpBase):
|
|||||||
|
|
||||||
|
|
||||||
class ExprGetter:
|
class ExprGetter:
|
||||||
def __init__(self, obj, ctx, *args, **kwargs):
|
def __init__(self, obj, ctx, args, kwargs, *, prepend_args=()):
|
||||||
self.obj = obj
|
self.obj = obj
|
||||||
self.ctx = ctx
|
self.ctx = ctx
|
||||||
self.args = args
|
self.args = args
|
||||||
|
self.prepend_args = prepend_args
|
||||||
self.kwargs = kwargs
|
self.kwargs = kwargs
|
||||||
|
|
||||||
def __call__(self, k, *, default=Empty):
|
def __call__(self, k, *, default=Empty):
|
||||||
obj = self.obj
|
obj = self.obj
|
||||||
result = (
|
if isinstance(k, str):
|
||||||
obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs)
|
result = obj.kwargs.get_eval(
|
||||||
if isinstance(k, str)
|
k, self.ctx, *self.args, default=default, **self.kwargs
|
||||||
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
|
)
|
||||||
)
|
elif isinstance(k, int):
|
||||||
|
pa = self.prepend_args
|
||||||
|
pa_len = len(pa)
|
||||||
|
result = (
|
||||||
|
pa[k]
|
||||||
|
if k < pa_len
|
||||||
|
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
|
||||||
|
)
|
||||||
if result is Empty:
|
if result is Empty:
|
||||||
raise KeyError(f"Unknown key {k!r}")
|
raise KeyError(f"Unknown key {k!r}")
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
class ExpMethodAp(ExpBase):
|
||||||
|
__slots__ = ("funap", "object_expression")
|
||||||
|
|
||||||
|
def __init__(self, object_expression, funap):
|
||||||
|
super().__init__()
|
||||||
|
self.object_expression = object_expression
|
||||||
|
self.funap = funap
|
||||||
|
|
||||||
|
def eval(self, handlers, *args, **kwargs):
|
||||||
|
object_value = self.object_expression.eval(handlers, *args, **kwargs)
|
||||||
|
type_name = type(object_value).__name__
|
||||||
|
handler_key = f"{type_name}::{self.funap.name}"
|
||||||
|
handler = handlers.get_handler(handler_key)
|
||||||
|
if handler is Empty:
|
||||||
|
raise KeyError(f"No handler for method call op: {handler_key!r}")
|
||||||
|
return handler(
|
||||||
|
self,
|
||||||
|
getter=ExprGetter(
|
||||||
|
self.funap, handlers, args, kwargs, prepend_args=(object_value,)
|
||||||
|
),
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
def clone(self, **kwargs):
|
||||||
|
return self.__class__(
|
||||||
|
object_expression=self.object_expression.clone(**kwargs),
|
||||||
|
funap=self.funap.clone(**kwargs),
|
||||||
|
)
|
||||||
|
|
||||||
|
__copy__ = clone
|
||||||
|
|
||||||
|
def __getattr__(self, k):
|
||||||
|
if k == "name":
|
||||||
|
return f"method::{self.funap.name}"
|
||||||
|
if k == "args":
|
||||||
|
return self.funap.args
|
||||||
|
if k == "kwargs":
|
||||||
|
return self.funap.kwargs
|
||||||
|
# This doesn't play well with deep copy.
|
||||||
|
# if hasattr(self.funap, k):
|
||||||
|
# return getattr(self.funap, k)
|
||||||
|
raise AttributeError(f"Can't get attribute {k}")
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"<METHAP:{self.object_expression}::{self.funap}>"
|
||||||
|
|
||||||
|
|
||||||
class ExpFunAp(ExpBase):
|
class ExpFunAp(ExpBase):
|
||||||
__slots__ = ("name", "args", "kwargs")
|
__slots__ = ("args", "kwargs", "name")
|
||||||
|
|
||||||
def __init__(self, name, args=None, kwargs=None):
|
def __init__(self, name, args=None, kwargs=None):
|
||||||
self.name = name
|
self.name = name
|
||||||
@@ -165,12 +230,14 @@ class ExpFunAp(ExpBase):
|
|||||||
handler = handlers.get_handler(self.name)
|
handler = handlers.get_handler(self.name)
|
||||||
if handler is Empty:
|
if handler is Empty:
|
||||||
raise KeyError(f"No handler for op: {self.name!r}")
|
raise KeyError(f"No handler for op: {self.name!r}")
|
||||||
return handler(
|
return handler(self, getter=ExprGetter(self, handlers, args, kwargs), **kwargs)
|
||||||
self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
def clone(self):
|
def clone(self, **kwargs):
|
||||||
return self.__class__(self.name, self.args.clone(), self.kwargs.clone())
|
return self.__class__(
|
||||||
|
self.name,
|
||||||
|
self.args.clone(**kwargs),
|
||||||
|
self.kwargs.clone(**kwargs),
|
||||||
|
)
|
||||||
|
|
||||||
def pretty_string(self, depth=0):
|
def pretty_string(self, depth=0):
|
||||||
pad = " " * (depth + 1) * 2
|
pad = " " * (depth + 1) * 2
|
||||||
@@ -202,12 +269,13 @@ class ExpBoundFunAp(ExpFunAp):
|
|||||||
|
|
||||||
__all__ = (
|
__all__ = (
|
||||||
"ExpBase",
|
"ExpBase",
|
||||||
"ExpOp",
|
|
||||||
"ExpBinOp",
|
"ExpBinOp",
|
||||||
"ExpSym",
|
"ExpBoundFunAp",
|
||||||
"ExpTuple",
|
|
||||||
"ExpKV",
|
|
||||||
"ExpDict",
|
"ExpDict",
|
||||||
"ExpFunAp",
|
"ExpFunAp",
|
||||||
"ExpBoundFunAp",
|
"ExpKV",
|
||||||
|
"ExpMethodAp",
|
||||||
|
"ExpOp",
|
||||||
|
"ExpSym",
|
||||||
|
"ExpTuple",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,12 +2,12 @@ import contextlib
|
|||||||
import functools
|
import functools
|
||||||
|
|
||||||
from ..latent import ImageBatch
|
from ..latent import ImageBatch
|
||||||
from .util import torch
|
|
||||||
from .types import Empty
|
from .types import Empty
|
||||||
|
from .util import torch
|
||||||
|
|
||||||
|
|
||||||
class Arg:
|
class Arg:
|
||||||
__slots__ = ("name", "default", "validator")
|
__slots__ = ("default", "name", "validator")
|
||||||
|
|
||||||
def __init__(self, name, default=Empty, *, validator=None):
|
def __init__(self, name, default=Empty, *, validator=None):
|
||||||
self.name = name
|
self.name = name
|
||||||
@@ -53,6 +53,23 @@ class Arg:
|
|||||||
name, default=default, validator=ValidateArg.validate_numscalar_sequence
|
name, default=default, validator=ValidateArg.validate_numscalar_sequence
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def numeric_sequence(cls, name, default=Empty):
|
||||||
|
return cls(
|
||||||
|
name, default=default, validator=ValidateArg.validate_numeric_sequence
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def numscalar_sequence_or_single(cls, name, default=Empty):
|
||||||
|
return cls.one_of(
|
||||||
|
name,
|
||||||
|
(
|
||||||
|
ValidateArg.validate_numscalar_sequence,
|
||||||
|
ValidateArg.validate_numeric_scalar,
|
||||||
|
),
|
||||||
|
default=default,
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def tensor_slice(cls, name, default=Empty):
|
def tensor_slice(cls, name, default=Empty):
|
||||||
return cls(name, default=default, validator=ValidateArg.validate_tensor_slice)
|
return cls(name, default=default, validator=ValidateArg.validate_tensor_slice)
|
||||||
@@ -86,8 +103,8 @@ class Arg:
|
|||||||
return cls(name, default=default, validator=ValidateArg.validate_boolean)
|
return cls(name, default=default, validator=ValidateArg.validate_boolean)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def present(cls, name):
|
def present(cls, name, default=Empty):
|
||||||
return cls(name, validator=ValidateArg.validate_passthrough)
|
return cls(name, default=default, validator=ValidateArg.validate_passthrough)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def one_of(cls, name, validators, *, default=Empty):
|
def one_of(cls, name, validators, *, default=Empty):
|
||||||
@@ -109,7 +126,7 @@ class ValidateError(Exception):
|
|||||||
|
|
||||||
|
|
||||||
class ValidateArg:
|
class ValidateArg:
|
||||||
__slots__ = ("valfuns", "groupfun", "kwargs", "kwargslist")
|
__slots__ = ("groupfun", "kwargs", "kwargslist", "valfuns")
|
||||||
|
|
||||||
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
|
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
|
||||||
if not isinstance(name, (list, tuple)):
|
if not isinstance(name, (list, tuple)):
|
||||||
@@ -224,6 +241,10 @@ class ValidateArg:
|
|||||||
idx, val, item_validator=cls.validate_numeric_scalar
|
idx, val, item_validator=cls.validate_numeric_scalar
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_numeric_sequence(cls, idx, val):
|
||||||
|
return cls.validate_sequence(idx, val, item_validator=cls.validate_numeric)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_tensor_slice(cls, idx, val):
|
def validate_tensor_slice(cls, idx, val):
|
||||||
return cls.validate_sequence(
|
return cls.validate_sequence(
|
||||||
@@ -236,6 +257,12 @@ class ValidateArg:
|
|||||||
raise ValidateError(f"Expected string argument at {idx}, got {type(val)}")
|
raise ValidateError(f"Expected string argument at {idx}, got {type(val)}")
|
||||||
return val
|
return val
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_dict(cls, idx, val):
|
||||||
|
if not isinstance(val, dict):
|
||||||
|
raise ValidateError(f"Expected dict argument at {idx}, got {type(val)}")
|
||||||
|
return val
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_boolean(cls, idx, val):
|
def validate_boolean(cls, idx, val):
|
||||||
if val is not True and val is not False:
|
if val is not True and val is not False:
|
||||||
|
|||||||
+456
-49
@@ -1,18 +1,23 @@
|
|||||||
import os
|
import os
|
||||||
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import PIL.Image as PILImage
|
import PIL.Image as PILImage
|
||||||
from functools import partial
|
import torch
|
||||||
|
|
||||||
from . import expression as expr
|
from . import expression as expr
|
||||||
from . import latent
|
from . import latent, unsafe_expression_whitelists
|
||||||
from . import unsafe_expression_whitelists
|
|
||||||
|
|
||||||
from .external import MODULES as EXT
|
from .external import MODULES as EXT
|
||||||
from .utils import scale_noise, resolve_value, quantile_normalize
|
from .latent import (
|
||||||
from .latent import OCSTAESD, ImageBatch, normalize_to_scale
|
OCSTAESD,
|
||||||
|
DimCorrelationConfig,
|
||||||
|
ImageBatch,
|
||||||
|
normalize_to_scale,
|
||||||
|
randomized_svd,
|
||||||
|
)
|
||||||
|
from .quantile_norm import quantile_normalize
|
||||||
|
from .utils import flip_tensor_range, resolve_value, scale_noise, softplus_soft_clamp
|
||||||
|
|
||||||
|
F = torch.nn.functional
|
||||||
|
|
||||||
ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None
|
ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None
|
||||||
ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None
|
ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None
|
||||||
@@ -42,6 +47,20 @@ def init_integrations(integrations):
|
|||||||
EXT.register_init_handler(init_integrations)
|
EXT.register_init_handler(init_integrations)
|
||||||
|
|
||||||
|
|
||||||
|
class UnaryTensorOpHandler(expr.BaseHandler):
|
||||||
|
input_validators = (expr.Arg.tensor("tensor"),)
|
||||||
|
|
||||||
|
def __init__(self, handler):
|
||||||
|
super().__init__()
|
||||||
|
self.handler = handler
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
(tensor,) = self.safe_get_all(obj, getter)
|
||||||
|
return self.handler(tensor)
|
||||||
|
|
||||||
|
validate_output = expr.Arg.tensor("output")
|
||||||
|
|
||||||
|
|
||||||
class NormHandler(expr.BaseHandler):
|
class NormHandler(expr.BaseHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor"),
|
expr.Arg.tensor("tensor"),
|
||||||
@@ -107,6 +126,20 @@ class ClampHandler(NormHandler):
|
|||||||
return torch.clamp(tensor, min=tmin, max=tmax)
|
return torch.clamp(tensor, min=tmin, max=tmax)
|
||||||
|
|
||||||
|
|
||||||
|
class SoftClampHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.numeric("min", 0.0),
|
||||||
|
expr.Arg.numeric("max", 1.0),
|
||||||
|
expr.Arg.numeric_scalar("stiffness", 1.0),
|
||||||
|
expr.Arg.boolean("safe", True),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, tmin, tmax, stiffness, safe = self.safe_get_all(obj, getter)
|
||||||
|
return softplus_soft_clamp(tensor, tmin, tmax, stiffness=stiffness, safe=safe)
|
||||||
|
|
||||||
|
|
||||||
class StackHandler(NormHandler):
|
class StackHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor),
|
expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor),
|
||||||
@@ -135,17 +168,66 @@ class ReshapeHandler(NormHandler):
|
|||||||
return torch.reshape(tensor.clone(), shape)
|
return torch.reshape(tensor.clone(), shape)
|
||||||
|
|
||||||
|
|
||||||
|
class SplitHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("chunk_size"),
|
||||||
|
expr.Arg.integer("dim"),
|
||||||
|
expr.Arg.boolean("pad_last", False),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, chunk_size, dim, pad_last = self.safe_get_all(obj, getter)
|
||||||
|
chunks = torch.split(tensor, chunk_size, dim=dim)
|
||||||
|
if not chunks or not pad_last or chunks[-1].shape == chunks[0].shape:
|
||||||
|
return chunks
|
||||||
|
dsize = chunks[0].shape[dim]
|
||||||
|
replacement_chunk = tensor.new_zeros(chunks[0].shape)
|
||||||
|
replacement_chunk[
|
||||||
|
tuple(
|
||||||
|
slice(None, dsize if d == dim else None) for d in replacement_chunk.ndim
|
||||||
|
)
|
||||||
|
] = chunks[-1]
|
||||||
|
return (*chunks[:-1], replacement_chunk)
|
||||||
|
|
||||||
|
|
||||||
|
class TrimHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.numscalar_sequence("shape"),
|
||||||
|
expr.Arg.boolean("flip", False),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, shape, flip = self.safe_get_all(obj, getter)
|
||||||
|
tshape = tensor.shape
|
||||||
|
if tensor.ndim != len(shape):
|
||||||
|
raise ValueError("Shape dimension count does not match tensor")
|
||||||
|
slices = tuple(
|
||||||
|
slice(None)
|
||||||
|
if tsize <= dsize
|
||||||
|
else (
|
||||||
|
slice(None, min(dsize, tsize))
|
||||||
|
if not flip
|
||||||
|
else slice(-min(dsize, tsize))
|
||||||
|
)
|
||||||
|
for dsize, tsize in zip(shape, tshape, strict=True)
|
||||||
|
)
|
||||||
|
return tensor[slices]
|
||||||
|
|
||||||
|
|
||||||
class IndexedCopyHandler(NormHandler):
|
class IndexedCopyHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor_dest"),
|
expr.Arg.tensor("tensor_dest"),
|
||||||
expr.Arg.tensor("tensor_src"),
|
expr.Arg.tensor("tensor_src"),
|
||||||
expr.Arg.tensor_slice("slice"),
|
expr.Arg.tensor_slice("slice"),
|
||||||
|
expr.Arg.boolean("slice_src", default=True),
|
||||||
)
|
)
|
||||||
|
|
||||||
def handle(self, obj, getter):
|
def handle(self, obj, getter):
|
||||||
tensor1, tensor2, tensor_slice = self.safe_get_all(obj, getter)
|
tensor1, tensor2, tensor_slice, slice_src = self.safe_get_all(obj, getter)
|
||||||
result = tensor1.clone()
|
result = tensor1.clone()
|
||||||
result[tensor_slice] = tensor2[tensor_slice]
|
result[tensor_slice] = tensor2[tensor_slice] if slice_src else tensor2
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
@@ -167,40 +249,36 @@ class NewLikeHandler(NormHandler):
|
|||||||
class MeanHandler(NormHandler):
|
class MeanHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor"),
|
expr.Arg.tensor("tensor"),
|
||||||
expr.Arg.numscalar_sequence("dim", (-3, -2, -1)),
|
expr.Arg.numscalar_sequence_or_single("dim", (-3, -2, -1)),
|
||||||
)
|
)
|
||||||
|
|
||||||
def handle(self, obj, getter):
|
def handle(self, obj, getter):
|
||||||
tensor, dim = self.safe_get_all(obj, getter)
|
tensor, dim = self.safe_get_all(obj, getter)
|
||||||
|
dim = dim if isinstance(dim, tuple) else (dim,)
|
||||||
return tensor.mean(keepdim=True, dim=dim)
|
return tensor.mean(keepdim=True, dim=dim)
|
||||||
|
|
||||||
|
|
||||||
class StdHandler(NormHandler):
|
class StdHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor"),
|
expr.Arg.tensor("tensor"),
|
||||||
expr.Arg.numscalar_sequence("dim", (-3, -2, -1)),
|
expr.Arg.numscalar_sequence_or_single("dim", (-3, -2, -1)),
|
||||||
|
expr.Arg.numeric("eps", default=1e-07),
|
||||||
)
|
)
|
||||||
|
|
||||||
def handle(self, obj, getter):
|
def handle(self, obj, getter):
|
||||||
tensor, dim = self.safe_get_all(obj, getter)
|
tensor, dim, eps = self.safe_get_all(obj, getter)
|
||||||
return tensor.std(keepdim=True, dim=dim)
|
dim = dim if isinstance(dim, tuple) else (dim,)
|
||||||
|
std = tensor.std(keepdim=True, dim=dim)
|
||||||
|
if eps != 0:
|
||||||
|
std = std.clamp_min_(eps)
|
||||||
|
return std
|
||||||
|
|
||||||
|
|
||||||
class RollHandler(NormHandler):
|
class RollHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor"),
|
expr.Arg.tensor("tensor"),
|
||||||
expr.Arg.numeric_scalar("amount", 0.5),
|
expr.Arg.numeric_scalar("amount", 0.5),
|
||||||
expr.Arg.one_of(
|
expr.Arg.numscalar_sequence_or_single("dim", -2),
|
||||||
"dim",
|
|
||||||
(
|
|
||||||
expr.ValidateArg.validate_integer,
|
|
||||||
partial(
|
|
||||||
expr.ValidateArg.validate_sequence,
|
|
||||||
item_validator=expr.ValidateArg.validate_integer,
|
|
||||||
),
|
|
||||||
),
|
|
||||||
default=-2,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def handle(self, obj, getter):
|
def handle(self, obj, getter):
|
||||||
@@ -275,20 +353,166 @@ class NewFullHandler(NormHandler):
|
|||||||
return tensor.new_full(shape, value)
|
return tensor.new_full(shape, value)
|
||||||
|
|
||||||
|
|
||||||
|
class InvertRangeHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("dim"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, dim = self.safe_get_all(obj, getter)
|
||||||
|
if dim < 0:
|
||||||
|
dim += tensor.ndim
|
||||||
|
if dim < 0 or dim >= tensor.ndim:
|
||||||
|
raise ValueError(
|
||||||
|
f"Dimension out of range, wanted {dim}, tensor has {tensor.ndim} dimension(s)"
|
||||||
|
)
|
||||||
|
return flip_tensor_range(tensor, dim=dim)
|
||||||
|
|
||||||
|
|
||||||
|
class MinHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("dim"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, dim = self.safe_get_all(obj, getter)
|
||||||
|
return tensor.min(dim=dim, keepdim=True).values
|
||||||
|
|
||||||
|
|
||||||
|
class MaxHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("dim"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, dim = self.safe_get_all(obj, getter)
|
||||||
|
return tensor.max(dim=dim, keepdim=True).values
|
||||||
|
|
||||||
|
|
||||||
|
class CumSumHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("dim"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, dim = self.safe_get_all(obj, getter)
|
||||||
|
return tensor.cumsum(dim=dim)
|
||||||
|
|
||||||
|
|
||||||
|
class MinimumHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor1"),
|
||||||
|
expr.Arg.tensor("tensor2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor1, tensor2 = self.safe_get_all(obj, getter)
|
||||||
|
return tensor1.minimum(tensor2)
|
||||||
|
|
||||||
|
|
||||||
|
class MaximumHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor1"),
|
||||||
|
expr.Arg.tensor("tensor2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor1, tensor2 = self.safe_get_all(obj, getter)
|
||||||
|
return tensor1.maximum(tensor2)
|
||||||
|
|
||||||
|
|
||||||
|
class MoveDimHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("from_dim"),
|
||||||
|
expr.Arg.integer("to_dim", -1),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, from_dim, to_dim = self.safe_get_all(obj, getter)
|
||||||
|
return tensor.movedim(from_dim, to_dim)
|
||||||
|
|
||||||
|
|
||||||
|
class PermuteHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.numscalar_sequence("perms"),
|
||||||
|
expr.Arg.boolean("reverse", default=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, perms, reverse = self.safe_get_all(obj, getter)
|
||||||
|
if not reverse:
|
||||||
|
return tensor.permute(tuple(perms))
|
||||||
|
inv_perms = torch.nn.utils.rnn.invert_permutation(
|
||||||
|
torch.tensor(perms, device="cpu"),
|
||||||
|
)
|
||||||
|
if inv_perms is None:
|
||||||
|
raise RuntimeError("Failed to calculate reverse permutation")
|
||||||
|
return tensor.permute(tuple(inv_perms.tolist()))
|
||||||
|
|
||||||
|
|
||||||
|
class FlattenHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.integer("start_dim", 0),
|
||||||
|
expr.Arg.integer("end_dim", -1),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, start_dim, end_dim = self.safe_get_all(obj, getter)
|
||||||
|
return tensor.flatten(start_dim=start_dim, end_dim=end_dim)
|
||||||
|
|
||||||
|
|
||||||
|
class MatmulHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor1"),
|
||||||
|
expr.Arg.tensor("tensor2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor1, tensor2 = self.safe_get_all(obj, getter)
|
||||||
|
return tensor1 @ tensor2
|
||||||
|
|
||||||
|
|
||||||
class BlendHandler(NormHandler):
|
class BlendHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor1"),
|
expr.Arg.tensor("tensor1"),
|
||||||
expr.Arg.tensor("tensor2"),
|
expr.Arg.tensor("tensor2"),
|
||||||
expr.Arg.numeric("scale", 0.5),
|
expr.Arg.numeric("scale", 0.5),
|
||||||
expr.Arg.string("mode", "lerp"),
|
expr.Arg.one_of(
|
||||||
|
"mode",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_string,
|
||||||
|
expr.ValidateArg.validate_dict,
|
||||||
|
),
|
||||||
|
default="lerp",
|
||||||
|
),
|
||||||
|
expr.Arg.one_of(
|
||||||
|
"blend_kwargs",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_none,
|
||||||
|
expr.ValidateArg.validate_dict,
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def handle(self, obj, getter):
|
def handle(self, obj, getter):
|
||||||
t1, t2, scale, mode = self.safe_get_all(obj, getter)
|
t1, t2, scale, mode, blend_kwargs = self.safe_get_all(obj, getter)
|
||||||
blend_handler = BLENDING_MODES.get(mode)
|
blend_handler = BLENDING_MODES.get(mode)
|
||||||
if not blend_handler:
|
if not blend_handler:
|
||||||
raise KeyError(f"Unknown blend mode {mode!r}")
|
raise KeyError(f"Unknown blend mode {mode!r}")
|
||||||
return blend_handler(t1, t2, scale)
|
return blend_handler(
|
||||||
|
t1,
|
||||||
|
t2,
|
||||||
|
scale,
|
||||||
|
**({} if blend_kwargs is None else blend_kwargs),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ContrastAdaptiveSharpeningHandler(NormHandler):
|
class ContrastAdaptiveSharpeningHandler(NormHandler):
|
||||||
@@ -348,7 +572,13 @@ class NoiseHandler(NormHandler):
|
|||||||
ctx.get_var(k, default=0.0)
|
ctx.get_var(k, default=0.0)
|
||||||
for k in ("sigma_min", "sigma_max", "sigma", "sigma_next")
|
for k in ("sigma_min", "sigma_max", "sigma", "sigma_next")
|
||||||
)
|
)
|
||||||
ns = latent.get_noise_sampler(typ, t, smin, smax, normalized=False)
|
ns = latent.get_noise_sampler(
|
||||||
|
typ,
|
||||||
|
t,
|
||||||
|
smin,
|
||||||
|
smax,
|
||||||
|
normalized=False,
|
||||||
|
)
|
||||||
return ns(s, sn)
|
return ns(s, sn)
|
||||||
|
|
||||||
|
|
||||||
@@ -360,6 +590,57 @@ class ShapeHandler(expr.BaseHandler):
|
|||||||
return expr.types.ExpTuple((*t.shape,))
|
return expr.types.ExpTuple((*t.shape,))
|
||||||
|
|
||||||
|
|
||||||
|
class NumelHandler(expr.BaseHandler):
|
||||||
|
input_validators = (expr.Arg.tensor("tensor"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
t = self.safe_get("tensor", obj, getter)
|
||||||
|
return t.numel()
|
||||||
|
|
||||||
|
|
||||||
|
class QRHandler(expr.BaseHandler):
|
||||||
|
input_validators = (expr.Arg.tensor("tensor"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor = self.safe_get("tensor", obj, getter)
|
||||||
|
return tuple(torch.linalg.qr(tensor))
|
||||||
|
|
||||||
|
|
||||||
|
class SVDHandler(expr.BaseHandler):
|
||||||
|
input_validators = (expr.Arg.tensor("tensor"),)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor = self.safe_get("tensor", obj, getter)
|
||||||
|
return tuple(torch.linalg.svd(tensor, full_matrices=False))
|
||||||
|
|
||||||
|
|
||||||
|
class RandomizedSVDHandler(expr.BaseHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor"),
|
||||||
|
expr.Arg.one_of(
|
||||||
|
"base",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_none,
|
||||||
|
expr.ValidateArg.validate_tensor,
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
),
|
||||||
|
expr.Arg.integer("n_iter", default=6),
|
||||||
|
expr.Arg.one_of(
|
||||||
|
"rank",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_none,
|
||||||
|
expr.ValidateArg.validate_integer,
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
tensor, base, n_iter, rank = self.safe_get_all(obj, getter)
|
||||||
|
return randomized_svd(tensor, y=base, n_iter=n_iter, rank=rank)
|
||||||
|
|
||||||
|
|
||||||
class GaussianBlur2DHandler(NormHandler):
|
class GaussianBlur2DHandler(NormHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("tensor"),
|
expr.Arg.tensor("tensor"),
|
||||||
@@ -385,6 +666,40 @@ class SNFGuidanceHandler(NormHandler):
|
|||||||
return latent.snf_guidance(*self.safe_get_all(obj, getter))
|
return latent.snf_guidance(*self.safe_get_all(obj, getter))
|
||||||
|
|
||||||
|
|
||||||
|
class CorrelateHandler(NormHandler):
|
||||||
|
input_validators = (
|
||||||
|
expr.Arg.tensor("tensor1"),
|
||||||
|
expr.Arg.one_of(
|
||||||
|
"tensor2",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_none,
|
||||||
|
expr.ValidateArg.validate_tensor,
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
),
|
||||||
|
expr.Arg.one_of(
|
||||||
|
"kwargs",
|
||||||
|
(
|
||||||
|
expr.ValidateArg.validate_none,
|
||||||
|
expr.ValidateArg.validate_dict,
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
t1, t2, kwargs = self.safe_get_all(obj, getter)
|
||||||
|
kwargs = kwargs.copy() if kwargs is not None else {}
|
||||||
|
if "leave" not in kwargs:
|
||||||
|
kwargs["leave"] = True
|
||||||
|
return (
|
||||||
|
DimCorrelationConfig.build(**kwargs)
|
||||||
|
.get_correlation_order(x=t1, ref=t2)
|
||||||
|
.reorder(t1)
|
||||||
|
.clone()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RGBLatentHandler(expr.BaseHandler):
|
class RGBLatentHandler(expr.BaseHandler):
|
||||||
input_validators = (
|
input_validators = (
|
||||||
expr.Arg.tensor("reference"),
|
expr.Arg.tensor("reference"),
|
||||||
@@ -626,42 +941,134 @@ class ScaleNNLatentUpscaleHandler(expr.BaseHandler):
|
|||||||
return latent.scale_nnlatentupscale(mode, tensor, scale)
|
return latent.scale_nnlatentupscale(mode, tensor, scale)
|
||||||
|
|
||||||
|
|
||||||
TENSOR_OP_HANDLERS = {
|
class ForkRngHandler(expr.BaseHandler):
|
||||||
"t_norm": NormHandler(),
|
input_validators = (
|
||||||
"t_quantilenorm": QuantileNormHandler(),
|
expr.Arg.present("expression"),
|
||||||
"t_normtoscale": NormToScaleHandler(),
|
expr.Arg.one_of(
|
||||||
"t_normalize_to_scale": NormToScaleHandler(),
|
"seed",
|
||||||
"t_reshape": ReshapeHandler(),
|
(
|
||||||
"t_clamp": ClampHandler(),
|
expr.ValidateArg.validate_none,
|
||||||
"t_cat": CatHandler(),
|
expr.ValidateArg.validate_integer,
|
||||||
"t_stack": StackHandler(),
|
),
|
||||||
"t_indexed_copy": IndexedCopyHandler(),
|
default=None,
|
||||||
"t_new_like": NewLikeHandler(),
|
),
|
||||||
"t_mean": MeanHandler(),
|
expr.Arg.boolean("enabled", default=True),
|
||||||
"t_std": StdHandler(),
|
)
|
||||||
|
|
||||||
|
def handle(self, obj, getter):
|
||||||
|
enabled = bool(self.safe_get("enabled", obj, getter))
|
||||||
|
seed = self.safe_get("seed", obj, getter)
|
||||||
|
with torch.random.fork_rng(enabled=enabled):
|
||||||
|
if enabled and seed is not None:
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
return self.safe_get("expression", obj, getter)
|
||||||
|
|
||||||
|
|
||||||
|
SKIP_TORCH_OPS = frozenset(
|
||||||
|
(
|
||||||
|
"fork_rng",
|
||||||
|
"t_cat",
|
||||||
|
"t_stack",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
TORCH_OP_HANDLERS = {
|
||||||
|
"fork_rng": ForkRngHandler(),
|
||||||
"t_blend": BlendHandler(),
|
"t_blend": BlendHandler(),
|
||||||
"t_roll": RollHandler(),
|
"t_cat": CatHandler(),
|
||||||
"t_flip": FlipHandler(),
|
"t_clamp": ClampHandler(),
|
||||||
|
"t_soft_clamp": SoftClampHandler(),
|
||||||
"t_clone": CloneHandler(),
|
"t_clone": CloneHandler(),
|
||||||
"t_newfull": NewFullHandler(),
|
|
||||||
"t_copysign": CopySignHandler(),
|
|
||||||
"t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(),
|
"t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(),
|
||||||
"t_scale": ScaleHandler(),
|
"t_copysign": CopySignHandler(),
|
||||||
"t_noise": NoiseHandler(),
|
"t_correlate": CorrelateHandler(),
|
||||||
"t_shape": ShapeHandler(),
|
"t_cumsum": CumSumHandler(),
|
||||||
|
"t_flatten": FlattenHandler(),
|
||||||
|
"t_flip": FlipHandler(),
|
||||||
"t_gaussianblur2d": GaussianBlur2DHandler(),
|
"t_gaussianblur2d": GaussianBlur2DHandler(),
|
||||||
|
"t_indexed_copy": IndexedCopyHandler(),
|
||||||
|
"t_invert_range": InvertRangeHandler(),
|
||||||
|
"t_matmul": MatmulHandler(),
|
||||||
|
"t_max": MaxHandler(),
|
||||||
|
"t_maximum": MaximumHandler(),
|
||||||
|
"t_mean": MeanHandler(),
|
||||||
|
"t_min": MinHandler(),
|
||||||
|
"t_minimum": MinimumHandler(),
|
||||||
|
"t_movedim": MoveDimHandler(),
|
||||||
|
"t_new_like": NewLikeHandler(),
|
||||||
|
"t_newfull": NewFullHandler(),
|
||||||
|
"t_noise": NoiseHandler(),
|
||||||
|
"t_norm": NormHandler(),
|
||||||
|
"t_normalize_to_scale": NormToScaleHandler(),
|
||||||
|
"t_normtoscale": NormToScaleHandler(),
|
||||||
|
"t_numel": NumelHandler(),
|
||||||
|
"t_quantilenorm": QuantileNormHandler(),
|
||||||
|
"t_reshape": ReshapeHandler(),
|
||||||
"t_rgb_latent": RGBLatentHandler(),
|
"t_rgb_latent": RGBLatentHandler(),
|
||||||
|
"t_roll": RollHandler(),
|
||||||
|
"t_scale": ScaleHandler(),
|
||||||
|
"t_shape": ShapeHandler(),
|
||||||
|
"t_permute": PermuteHandler(),
|
||||||
|
"t_qr": QRHandler(),
|
||||||
|
"t_svd": SVDHandler(),
|
||||||
|
"t_randomized_svd": RandomizedSVDHandler(),
|
||||||
"t_snf_guidance": SNFGuidanceHandler(),
|
"t_snf_guidance": SNFGuidanceHandler(),
|
||||||
|
"t_split": SplitHandler(),
|
||||||
|
"t_stack": StackHandler(),
|
||||||
|
"t_std": StdHandler(),
|
||||||
"t_taesd_decode": TAESDDecodeHandler(),
|
"t_taesd_decode": TAESDDecodeHandler(),
|
||||||
|
"t_trim": TrimHandler(),
|
||||||
"unsafe_tensor_method": UnsafeTorchTensorMethodHandler(),
|
"unsafe_tensor_method": UnsafeTorchTensorMethodHandler(),
|
||||||
"unsafe_torch": UnsafeTorchHandler(),
|
"unsafe_torch": UnsafeTorchHandler(),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
TORCH_UOP_HANDLERS = {
|
||||||
|
f"t_{k}": UnaryTensorOpHandler(handler=h)
|
||||||
|
for k, h in (
|
||||||
|
("abs", torch.abs),
|
||||||
|
("acos", torch.acos),
|
||||||
|
("acosh", torch.acosh),
|
||||||
|
("atan", torch.atan),
|
||||||
|
("atanh", torch.atanh),
|
||||||
|
("ceil", torch.ceil),
|
||||||
|
("clone", torch.clone),
|
||||||
|
("cos", torch.cos),
|
||||||
|
("erf", torch.erf),
|
||||||
|
("erfinv", torch.erfinv),
|
||||||
|
("exp", torch.exp),
|
||||||
|
("expm1", torch.expm1),
|
||||||
|
("floor", torch.floor),
|
||||||
|
("frac", torch.frac),
|
||||||
|
("gelu", F.gelu),
|
||||||
|
("log", torch.log),
|
||||||
|
("log1p", torch.log1p),
|
||||||
|
("reciprocal", torch.reciprocal),
|
||||||
|
("relu", F.relu),
|
||||||
|
("remainder", torch.remainder),
|
||||||
|
("rsqrt", torch.rsqrt),
|
||||||
|
("sigmoid", torch.sigmoid),
|
||||||
|
("sign", torch.sign),
|
||||||
|
("sin", torch.sin),
|
||||||
|
("tan", torch.tan),
|
||||||
|
("tanh", torch.tanh),
|
||||||
|
("trunc", torch.trunc),
|
||||||
|
("diag_embed", torch.diag_embed),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
TORCH_OP_HANDLERS |= TORCH_UOP_HANDLERS # ty: ignore[unsupported-operator]
|
||||||
|
|
||||||
|
TORCH_OP_HANDLERS |= {
|
||||||
|
f"Tensor::{k.removeprefix('t_')}": v
|
||||||
|
for k, v in TORCH_OP_HANDLERS.items()
|
||||||
|
if k not in SKIP_TORCH_OPS
|
||||||
|
}
|
||||||
|
|
||||||
IMAGE_OP_HANDLERS = {
|
IMAGE_OP_HANDLERS = {
|
||||||
"img_taesd_encode": TAESDEncodeHandler(),
|
"img_taesd_encode": TAESDEncodeHandler(),
|
||||||
"img_shape": ImgShapeHandler(),
|
"img_shape": ImgShapeHandler(),
|
||||||
"img_pil_resize": ImgPILResizeHandler(),
|
"img_pil_resize": ImgPILResizeHandler(),
|
||||||
}
|
}
|
||||||
|
|
||||||
HANDLERS |= TENSOR_OP_HANDLERS
|
HANDLERS |= TORCH_OP_HANDLERS
|
||||||
HANDLERS |= IMAGE_OP_HANDLERS
|
HANDLERS |= IMAGE_OP_HANDLERS
|
||||||
|
|||||||
+316
-9
@@ -1,13 +1,13 @@
|
|||||||
import numpy as np
|
from typing import Any, NamedTuple, Self
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import latent_preview
|
import latent_preview
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from comfy import latent_formats
|
||||||
from comfy.taesd.taesd import TAESD
|
from comfy.taesd.taesd import TAESD
|
||||||
from comfy.utils import bislerp
|
from comfy.utils import bislerp
|
||||||
from comfy import latent_formats
|
|
||||||
|
|
||||||
from .external import MODULES as EXT
|
from .external import MODULES as EXT
|
||||||
|
|
||||||
@@ -42,7 +42,7 @@ def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)):
|
|||||||
# Improvements by https://github.com/Clybius
|
# Improvements by https://github.com/Clybius
|
||||||
# The following is modified to work with latent images of ~0 mean from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/tree/main.
|
# The following is modified to work with latent images of ~0 mean from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/tree/main.
|
||||||
# The algorithm is directly implemented from FidelityFX's source code that can be found here: https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h.
|
# The algorithm is directly implemented from FidelityFX's source code that can be found here: https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h.
|
||||||
def contrast_adaptive_sharpening( # noqa: PLR0914
|
def contrast_adaptive_sharpening(
|
||||||
x,
|
x,
|
||||||
amount=0.8,
|
amount=0.8,
|
||||||
*,
|
*,
|
||||||
@@ -228,7 +228,7 @@ def scale_samples(
|
|||||||
|
|
||||||
def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): # noqa: F811
|
def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): # noqa: F811
|
||||||
if noise_type != "gaussian":
|
if noise_type != "gaussian":
|
||||||
raise ValueError("Only gaussian noise supported")
|
raise ValueError("Only gaussian noise supported unless you have ComfyUI-sonar")
|
||||||
return lambda _s, _sn: torch.randn_like(x)
|
return lambda _s, _sn: torch.randn_like(x)
|
||||||
|
|
||||||
|
|
||||||
@@ -324,7 +324,7 @@ class OCSLatentFormat:
|
|||||||
|
|
||||||
def latent_to_rgb(self, latent: torch.Tensor) -> torch.Tensor:
|
def latent_to_rgb(self, latent: torch.Tensor) -> torch.Tensor:
|
||||||
# NCHW -> NHWC
|
# NCHW -> NHWC
|
||||||
if self.latent_factors is None:
|
if self.rgb_factors is None:
|
||||||
raise ValueError("No RGB factors for latent type!")
|
raise ValueError("No RGB factors for latent type!")
|
||||||
return torch.nn.functional.linear(
|
return torch.nn.functional.linear(
|
||||||
latent.movedim(1, -1), self.rgb_factors, bias=self.rgb_factors_bias
|
latent.movedim(1, -1), self.rgb_factors, bias=self.rgb_factors_bias
|
||||||
@@ -332,8 +332,315 @@ class OCSLatentFormat:
|
|||||||
|
|
||||||
def rgb_to_latent(self, img: torch.Tensor) -> torch.Tensor:
|
def rgb_to_latent(self, img: torch.Tensor) -> torch.Tensor:
|
||||||
# NHWC
|
# NHWC
|
||||||
if self.latent_factors is None:
|
if self.rgb_factors is None:
|
||||||
raise ValueError("No RGB factors for latent type!")
|
raise ValueError("No RGB factors for latent type!")
|
||||||
if self.rgb_factors_bias is not None:
|
if self.rgb_factors_bias is not None:
|
||||||
img = img - self.rgb_factors_bias
|
img = img - self.rgb_factors_bias
|
||||||
return torch.nn.functional.linear(img, self.rgb_factors_inv)
|
return torch.nn.functional.linear(img, self.rgb_factors_inv)
|
||||||
|
|
||||||
|
|
||||||
|
def randomized_svd(
|
||||||
|
m: torch.Tensor,
|
||||||
|
*,
|
||||||
|
rank: int | None = None,
|
||||||
|
n_iter: int = 6,
|
||||||
|
ortho_interval: int = 3,
|
||||||
|
oversample: int = 10,
|
||||||
|
noise_sampler: Callable | None = None,
|
||||||
|
y: torch.Tensor | None = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
n, c = m.shape[-2], m.shape[-1]
|
||||||
|
if rank is None:
|
||||||
|
rank = n
|
||||||
|
if y is None:
|
||||||
|
if noise_sampler is None:
|
||||||
|
noise_sampler = torch.randn
|
||||||
|
k = min(rank + oversample, n, c)
|
||||||
|
y_shape = (*m.shape[:-2], c, k)
|
||||||
|
y = noise_sampler(y_shape, device=m.device, dtype=m.dtype)
|
||||||
|
elif y.shape == m.shape:
|
||||||
|
y = y.mT
|
||||||
|
elif y.shape != m.mT.shape:
|
||||||
|
raise ValueError("Bad initial y shape")
|
||||||
|
y = m @ y
|
||||||
|
|
||||||
|
ortho = False
|
||||||
|
for idx in range(n_iter):
|
||||||
|
y = m @ (m.mT @ y)
|
||||||
|
ortho = ortho_interval > 0 and (idx % ortho_interval) == 0 and n_iter - idx != 2
|
||||||
|
if ortho:
|
||||||
|
y = torch.linalg.qr(y)[0]
|
||||||
|
|
||||||
|
q = y if ortho else torch.linalg.qr(y)[0]
|
||||||
|
u, s, vh = torch.linalg.svd(q.mT @ m, full_matrices=False)
|
||||||
|
u = q @ u
|
||||||
|
if rank < n:
|
||||||
|
return u[..., :rank], s[..., :rank], vh[..., :rank, :]
|
||||||
|
return u, s, vh
|
||||||
|
|
||||||
|
|
||||||
|
class DimCorrelationOrder(NamedTuple):
|
||||||
|
perm: torch.Tensor
|
||||||
|
dim: int
|
||||||
|
leave: bool = False
|
||||||
|
|
||||||
|
def reorder(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
invert: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
dim, perm = self.dim, self.perm
|
||||||
|
if invert:
|
||||||
|
perm = perm.argsort(dim=-1)
|
||||||
|
if dim < 0:
|
||||||
|
dim = x.ndim + self.dim
|
||||||
|
shape = [1] * x.ndim
|
||||||
|
if dim != 0:
|
||||||
|
shape[0] = x.shape[0]
|
||||||
|
shape[dim] = x.shape[dim]
|
||||||
|
perm = perm.view(*shape).expand_as(x)
|
||||||
|
return x.gather(dim=dim, index=perm)
|
||||||
|
|
||||||
|
|
||||||
|
class DimCorrelationConfig(NamedTuple):
|
||||||
|
dim: int = 1
|
||||||
|
flip: bool = False
|
||||||
|
cross: bool = False
|
||||||
|
leave: bool = False
|
||||||
|
preserve_first: bool = False
|
||||||
|
center_strength: float = 1.0
|
||||||
|
center_dim: int = -1
|
||||||
|
# None - disabled, otherwise controls whether abs occurs before or after centering.
|
||||||
|
abs_before: bool | None = None
|
||||||
|
# 0 - disabled, positive value - enabled, negative value - enabled with sign flipped.
|
||||||
|
fix_sign: int = 1
|
||||||
|
align_to_peak: bool = False
|
||||||
|
expansion_factor: int = 1 # NYI
|
||||||
|
# Only applies to cross mode. One of: svd, randomized_svd
|
||||||
|
decomp_mode: str = "svd"
|
||||||
|
low_rank: int = 0
|
||||||
|
low_rank_niter: int = 6
|
||||||
|
work_dtype: torch.dtype | None = None
|
||||||
|
pc: int = 0
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def build(cls, **kwargs: Any) -> Self:
|
||||||
|
wd = kwargs.get("work_dtype")
|
||||||
|
if isinstance(wd, str):
|
||||||
|
dtype_map = {
|
||||||
|
"float64": torch.float64,
|
||||||
|
"float32": torch.float32,
|
||||||
|
"float16": torch.float16,
|
||||||
|
"bfloat16": torch.bfloat16,
|
||||||
|
}
|
||||||
|
wd = dtype_map.get(wd)
|
||||||
|
if wd is None:
|
||||||
|
raise ValueError("Bad dtype")
|
||||||
|
kwargs["work_dtype"] = wd
|
||||||
|
if kwargs.get("decomp_mode") not in {None, "svd", "randomized_svd"}:
|
||||||
|
raise ValueError(
|
||||||
|
"Bad decomp mode, must be unset or one of: svd, randomized_svd",
|
||||||
|
)
|
||||||
|
fs = frozenset(cls._fields)
|
||||||
|
kwargs = {k: v for k, v in kwargs.items() if k in fs}
|
||||||
|
return cls(**kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_str(cls, s: str) -> Self:
|
||||||
|
parts = tuple(p.strip() for p in s.split(":", 3))
|
||||||
|
plen = len(parts)
|
||||||
|
if plen > 2:
|
||||||
|
raise ValueError("Dim correlations only support up to two parts.")
|
||||||
|
dim = int(parts[0])
|
||||||
|
result = cls(dim=dim)
|
||||||
|
if plen < 2 or not parts[1]:
|
||||||
|
return result
|
||||||
|
p1 = parts[1]
|
||||||
|
p1len = len(p1)
|
||||||
|
offs = 0
|
||||||
|
while offs < p1len:
|
||||||
|
pflag = p1[offs]
|
||||||
|
offs += 1
|
||||||
|
if pflag == "f":
|
||||||
|
result = result._replace(flip=True)
|
||||||
|
elif pflag == "x":
|
||||||
|
result = result._replace(cross=True)
|
||||||
|
elif pflag == "l":
|
||||||
|
result = result._replace(leave=True)
|
||||||
|
elif pflag == "u":
|
||||||
|
result = result._replace(center_strength=0.0)
|
||||||
|
elif pflag == "c":
|
||||||
|
result = result._replace(center_dim=-2)
|
||||||
|
elif pflag in "aA":
|
||||||
|
result = result._replace(abs_before=pflag == "a")
|
||||||
|
elif pflag == "s":
|
||||||
|
result = result._replace(fix_sign=0)
|
||||||
|
elif pflag == "S":
|
||||||
|
result = result._replace(fix_sign=-1)
|
||||||
|
elif pflag == "p":
|
||||||
|
result = result._replace(align_to_peak=True)
|
||||||
|
elif pflag == "i":
|
||||||
|
result = result._replace(preserve_first=True)
|
||||||
|
else:
|
||||||
|
offs -= 1
|
||||||
|
break
|
||||||
|
p1 = p1[offs:].strip()
|
||||||
|
pc = int(p1) if p1 else 0
|
||||||
|
return result._replace(pc=pc)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _fix_sign_ambiguity(
|
||||||
|
cls,
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
ref: torch.Tensor | None = None,
|
||||||
|
dim: int = -1,
|
||||||
|
in_place: bool = True,
|
||||||
|
neg: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if ref is not None:
|
||||||
|
x = cls._fix_sign_ambiguity(x, dim=dim, in_place=in_place)
|
||||||
|
else:
|
||||||
|
ref = x
|
||||||
|
signs = ref.gather(dim, ref.abs().argmax(dim=dim, keepdim=True)).sign_()
|
||||||
|
signs = signs.masked_fill_(signs == 0, 1.0)
|
||||||
|
if neg:
|
||||||
|
signs = signs.neg_()
|
||||||
|
return x.mul_(signs) if in_place else x * signs
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _preserve_first_index(perm: torch.Tensor) -> torch.Tensor:
|
||||||
|
c = perm.shape[-1]
|
||||||
|
shift = (perm == 0).to(dtype=torch.int64).argmax(dim=-1, keepdim=True)
|
||||||
|
shift = torch.arange(c, device=perm.device, dtype=shift.dtype) + shift
|
||||||
|
shift %= c
|
||||||
|
return perm.gather(dim=-1, index=shift)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _align_ref(
|
||||||
|
*,
|
||||||
|
x: torch.Tensor,
|
||||||
|
ref: torch.Tensor,
|
||||||
|
skip_dim: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if skip_dim < 0:
|
||||||
|
skip_dim = x.ndim + skip_dim
|
||||||
|
reps = tuple(
|
||||||
|
None if szr == 0 or szx % szr != 0 else (d, szx // szr)
|
||||||
|
for d, (szx, szr) in enumerate(zip(x.shape, ref.shape, strict=True))
|
||||||
|
if d != skip_dim and szx != szr
|
||||||
|
)
|
||||||
|
if not all(reps):
|
||||||
|
raise ValueError("Bad shape")
|
||||||
|
for d, r in reps:
|
||||||
|
ref = ref.repeat_interleave(r, d)
|
||||||
|
return ref
|
||||||
|
|
||||||
|
def _get_correlation_order(
|
||||||
|
self,
|
||||||
|
cov: torch.Tensor,
|
||||||
|
*,
|
||||||
|
pc_idx: int = 0,
|
||||||
|
cross_mode: bool = False,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if cross_mode:
|
||||||
|
if self.decomp_mode == "randomized_svd":
|
||||||
|
pcs = randomized_svd(cov, n_iter=self.low_rank_niter, **kwargs)[0]
|
||||||
|
elif self.low_rank < 1:
|
||||||
|
pcs = torch.linalg.svd(cov, full_matrices=False).U
|
||||||
|
else:
|
||||||
|
pcs = torch.svd_lowrank(
|
||||||
|
cov,
|
||||||
|
q=self.low_rank,
|
||||||
|
niter=self.low_rank_niter,
|
||||||
|
)[0]
|
||||||
|
n_pcs = pcs.shape[-1]
|
||||||
|
if pc_idx < 0:
|
||||||
|
pc_idx = n_pcs + pc_idx
|
||||||
|
else:
|
||||||
|
pcs = torch.linalg.eigh(cov).eigenvectors
|
||||||
|
n_pcs = pcs.shape[-1]
|
||||||
|
pc_idx = n_pcs - pc_idx - 1 if pc_idx >= 0 else pc_idx + 1
|
||||||
|
pc_idx = max(0, min(n_pcs - 1, pc_idx))
|
||||||
|
return pcs[..., pc_idx]
|
||||||
|
|
||||||
|
def _preprocess(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
dim: int,
|
||||||
|
allow_post_abs: bool = True,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
x_flat = (x.unsqueeze(0) if dim == 0 else x.movedim(dim, 1)).flatten(
|
||||||
|
start_dim=2
|
||||||
|
)
|
||||||
|
if self.abs_before is True:
|
||||||
|
x_flat = x_flat.abs()
|
||||||
|
if self.center_strength == 0:
|
||||||
|
return x_flat
|
||||||
|
xm = x_flat.mean(dim=self.center_dim, keepdim=True)
|
||||||
|
if self.center_strength != 1:
|
||||||
|
xm *= self.center_strength
|
||||||
|
x_flat = x_flat.sub_(xm) if self.abs_before else x_flat - xm
|
||||||
|
if allow_post_abs and self.abs_before is False:
|
||||||
|
x_flat = x_flat.abs_()
|
||||||
|
return x_flat
|
||||||
|
|
||||||
|
def get_correlation_order(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
ref: torch.Tensor | None = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> DimCorrelationOrder:
|
||||||
|
if self.work_dtype is not None:
|
||||||
|
if x.dtype != self.work_dtype:
|
||||||
|
x = x.to(dtype=self.work_dtype)
|
||||||
|
if ref is not None and ref.dtype != self.work_dtype:
|
||||||
|
ref = ref.to(dtype=self.work_dtype)
|
||||||
|
dim, pc_idx = self.dim, self.pc
|
||||||
|
if dim < 0:
|
||||||
|
dim = x.ndim + dim
|
||||||
|
allow_post_abs = self.abs_before is not False or not self.align_to_peak
|
||||||
|
x_flat = self._preprocess(
|
||||||
|
x,
|
||||||
|
dim=dim,
|
||||||
|
allow_post_abs=ref is not None or allow_post_abs,
|
||||||
|
)
|
||||||
|
if ref is None or not self.cross:
|
||||||
|
# if ref is None:
|
||||||
|
y_flat = x_flat
|
||||||
|
if not allow_post_abs:
|
||||||
|
sign_ref = y_flat.clone()
|
||||||
|
y_flat = y_flat.abs_()
|
||||||
|
else:
|
||||||
|
sign_ref = y_flat if self.align_to_peak else None
|
||||||
|
else:
|
||||||
|
if ref.shape != x.shape:
|
||||||
|
ref = self._align_ref(x=x, ref=ref, skip_dim=dim)
|
||||||
|
y_flat = self._preprocess(ref, dim=dim, allow_post_abs=allow_post_abs)
|
||||||
|
if not allow_post_abs:
|
||||||
|
sign_ref = y_flat.clone()
|
||||||
|
y_flat = y_flat.abs_()
|
||||||
|
else:
|
||||||
|
sign_ref = y_flat if self.align_to_peak else None
|
||||||
|
pc = self._get_correlation_order(
|
||||||
|
x_flat @ y_flat.mT,
|
||||||
|
pc_idx=pc_idx,
|
||||||
|
cross_mode=ref is not None,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
if self.fix_sign:
|
||||||
|
pc = self._fix_sign_ambiguity(
|
||||||
|
pc,
|
||||||
|
neg=self.fix_sign < 0,
|
||||||
|
ref=None
|
||||||
|
if sign_ref is None
|
||||||
|
else sign_ref.flatten(start_dim=0 if dim == 0 else 1),
|
||||||
|
)
|
||||||
|
perm = pc.argsort(dim=-1, descending=self.flip)
|
||||||
|
if self.preserve_first:
|
||||||
|
perm = self._preserve_first_index(perm)
|
||||||
|
return DimCorrelationOrder(dim=dim, perm=perm, leave=self.leave)
|
||||||
|
|||||||
+5
-4
@@ -169,8 +169,9 @@ class OCSModel:
|
|||||||
self.extra_args = extra_args
|
self.extra_args = extra_args
|
||||||
self.cfg1_uncond_optimization = cfg1_uncond_optimization
|
self.cfg1_uncond_optimization = cfg1_uncond_optimization
|
||||||
self.cfg_scale_override = cfg_scale_override
|
self.cfg_scale_override = cfg_scale_override
|
||||||
|
self.model_sampling = model.inner_model.inner_model.model_sampling
|
||||||
self.is_rectified_flow = isinstance(
|
self.is_rectified_flow = isinstance(
|
||||||
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
|
self.model_sampling, comfy.model_sampling.CONST
|
||||||
)
|
)
|
||||||
self.latent_format = OCSLatentFormat(
|
self.latent_format = OCSLatentFormat(
|
||||||
x.device, model.inner_model.inner_model.latent_format
|
x.device, model.inner_model.inner_model.latent_format
|
||||||
@@ -212,9 +213,9 @@ class OCSModel:
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
||||||
|
|
||||||
@property
|
# @property
|
||||||
def model_sampling(self):
|
# def model_sampling(self):
|
||||||
return self.model.inner_model.inner_model.model_sampling
|
# return self.model.inner_model.inner_model.model_sampling
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def inner_cfg_scale(self) -> None | int | float:
|
def inner_cfg_scale(self) -> None | int | float:
|
||||||
|
|||||||
+343
-13
@@ -1,17 +1,17 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
import yaml
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import yaml
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from .external import MODULES, IntegratedNode
|
from .external import MODULES, IntegratedNode
|
||||||
|
from .filtering import Filter, FilterRefs, make_filter
|
||||||
from .restart import Restart
|
from .restart import Restart
|
||||||
from .sampling import composable_sampler
|
from .sampling import composable_sampler
|
||||||
from .step_samplers import STEP_SAMPLERS
|
from .step_samplers import STEP_SAMPLERS
|
||||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||||
from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups
|
from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups
|
||||||
from .filtering import make_filter
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from comfy_execution import validation as comfy_validation
|
from comfy_execution import validation as comfy_validation
|
||||||
@@ -27,15 +27,17 @@ except Exception as exc:
|
|||||||
f"** OCS: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}"
|
f"** OCS: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}"
|
||||||
)
|
)
|
||||||
|
|
||||||
PARAM_INPUT_TYPES = frozenset((
|
PARAM_INPUT_TYPES = frozenset(
|
||||||
"IMAGE",
|
(
|
||||||
"OCS_NOISE",
|
"IMAGE",
|
||||||
"SAMPLER",
|
"OCS_NOISE",
|
||||||
"SIGMAS",
|
"SAMPLER",
|
||||||
"SONAR_CUSTOM_NOISE",
|
"SIGMAS",
|
||||||
"UPSCALE_MODEL",
|
"SONAR_CUSTOM_NOISE",
|
||||||
"VAE",
|
"UPSCALE_MODEL",
|
||||||
))
|
"VAE",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
|
NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
|
||||||
|
|
||||||
@@ -778,6 +780,334 @@ class ApplyFilterImage(ApplyFilterLatent):
|
|||||||
return (result,)
|
return (result,)
|
||||||
|
|
||||||
|
|
||||||
|
class ExpressionFilteredLatentOperation:
|
||||||
|
EXTENDED_LATENT_OPERATION = True
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
ocs_filter: Filter,
|
||||||
|
latent_refs: dict[str, torch.Tensor] | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.filter = ocs_filter
|
||||||
|
self.latent_refs = latent_refs if latent_refs is not None else {}
|
||||||
|
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
latent: torch.Tensor,
|
||||||
|
*,
|
||||||
|
sigma: float | torch.Tensor | None = None,
|
||||||
|
**kwargs: dict,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
refs = FilterRefs(
|
||||||
|
kvs=kwargs
|
||||||
|
| {
|
||||||
|
"sigma": sigma.clone() if isinstance(sigma, torch.Tensor) else sigma,
|
||||||
|
"sigma_float": sigma.max().item()
|
||||||
|
if isinstance(sigma, torch.Tensor)
|
||||||
|
else sigma,
|
||||||
|
}
|
||||||
|
| {k: v.to(latent, copy=True) for k, v in self.latent_refs.items()}
|
||||||
|
)
|
||||||
|
return self.filter.apply(latent, refs=refs)
|
||||||
|
|
||||||
|
|
||||||
|
class ExpressionFilteredLatentOperationNode:
|
||||||
|
DESCRIPTION = "TBD"
|
||||||
|
|
||||||
|
FUNCTION = "go"
|
||||||
|
RETURN_TYPES = ("LATENT_OPERATION",)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
MODULES.initialize()
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"yaml_config": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"placeholder": """\
|
||||||
|
# YAML or JSON filter definition
|
||||||
|
""",
|
||||||
|
"multiline": True,
|
||||||
|
"dynamicPrompts": False,
|
||||||
|
"tooltip": "Enter your filter definition here. There is essentially no error handling.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"latent_ref_1_opt": ("LATENT",),
|
||||||
|
"latent_ref_2_opt": ("LATENT",),
|
||||||
|
"latent_ref_3_opt": ("LATENT",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
def go(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
yaml_config: str,
|
||||||
|
latent_ref_1_opt: dict | None = None,
|
||||||
|
latent_ref_2_opt: dict | None = None,
|
||||||
|
latent_ref_3_opt: dict | None = None,
|
||||||
|
) -> tuple:
|
||||||
|
config = yaml.safe_load(yaml_config)
|
||||||
|
if isinstance(config, str):
|
||||||
|
config = {"filter": {"final": config}}
|
||||||
|
elif not isinstance(config, dict) or "filter" not in config:
|
||||||
|
raise ValueError(
|
||||||
|
"Bad YAML config type (must be object) or missing filter key in config"
|
||||||
|
)
|
||||||
|
filter_def = config.get("filter")
|
||||||
|
if not isinstance(filter_def, dict):
|
||||||
|
raise TypeError("Bad type for filter definition, must be object")
|
||||||
|
latent_refs = {
|
||||||
|
k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True)
|
||||||
|
for k, v in (
|
||||||
|
("latent_ref_1", latent_ref_1_opt),
|
||||||
|
("latent_ref_2", latent_ref_2_opt),
|
||||||
|
("latent_ref_3", latent_ref_3_opt),
|
||||||
|
)
|
||||||
|
if v is not None
|
||||||
|
}
|
||||||
|
ocs_filter = make_filter(filter_def)
|
||||||
|
return (
|
||||||
|
ExpressionFilteredLatentOperation(
|
||||||
|
ocs_filter=ocs_filter, latent_refs=latent_refs
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ExpressionFilteredModelPatchNode:
|
||||||
|
DESCRIPTION = "TBD"
|
||||||
|
|
||||||
|
FUNCTION = "go"
|
||||||
|
RETURN_TYPES = ("MODEL",)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
MODULES.initialize()
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"model": ("MODEL",),
|
||||||
|
"patch_mode": (
|
||||||
|
("apply_model", "pre_cfg", "post_cfg", "cfg", "denoise_mask"),
|
||||||
|
{"default": "apply_model"},
|
||||||
|
),
|
||||||
|
"existing_patch_mode": (
|
||||||
|
("normal", "extract", "extract_split", "extract_sequence"),
|
||||||
|
{
|
||||||
|
"default": "normal",
|
||||||
|
"tooltip": "Modes:\n"
|
||||||
|
"normal: Replaces apply_model or cfg patches, appends for pre_cfg and post_cfg.\n"
|
||||||
|
"extract: Removes the existing patches and passes old_result with the output from existing patches.\n"
|
||||||
|
"extract_split: Same as extract except you'll get tuple of results for each existing patch (apply_model and cfg will always be length 1).\n"
|
||||||
|
"extract_sequence: Like extract_split except existing patches do not see each other's effects.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"yaml_config": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"placeholder": """\
|
||||||
|
# YAML or JSON filter definition
|
||||||
|
""",
|
||||||
|
"multiline": True,
|
||||||
|
"dynamicPrompts": False,
|
||||||
|
"tooltip": "Enter your filter definition here. There is essentially no error handling.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"latent_ref_1_opt": ("LATENT",),
|
||||||
|
"latent_ref_2_opt": ("LATENT",),
|
||||||
|
"latent_ref_3_opt": ("LATENT",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
def go(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
yaml_config: str,
|
||||||
|
model: object,
|
||||||
|
patch_mode: str,
|
||||||
|
existing_patch_mode: str,
|
||||||
|
latent_ref_1_opt: dict | None = None,
|
||||||
|
latent_ref_2_opt: dict | None = None,
|
||||||
|
latent_ref_3_opt: dict | None = None,
|
||||||
|
) -> tuple:
|
||||||
|
config = yaml.safe_load(yaml_config)
|
||||||
|
if isinstance(config, str):
|
||||||
|
config = {"filter": {"final": config}}
|
||||||
|
elif not isinstance(config, dict) or "filter" not in config:
|
||||||
|
raise ValueError(
|
||||||
|
"Bad YAML config type (must be object) or missing filter key in config"
|
||||||
|
)
|
||||||
|
filter_def = config.get("filter")
|
||||||
|
if not isinstance(filter_def, dict):
|
||||||
|
raise TypeError("Bad type for filter definition, must be object")
|
||||||
|
latent_refs = {
|
||||||
|
k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True)
|
||||||
|
for k, v in (
|
||||||
|
("latent_ref_1", latent_ref_1_opt),
|
||||||
|
("latent_ref_2", latent_ref_2_opt),
|
||||||
|
("latent_ref_3", latent_ref_3_opt),
|
||||||
|
)
|
||||||
|
if v is not None
|
||||||
|
}
|
||||||
|
ocs_filter = make_filter(filter_def)
|
||||||
|
model = model.clone()
|
||||||
|
mode_keys = {
|
||||||
|
"post_cfg": "sampler_post_cfg_function",
|
||||||
|
"pre_cfg": "sampler_pre_cfg_function",
|
||||||
|
"apply_model": "model_function_wrapper",
|
||||||
|
"cfg": "sampler_cfg_function",
|
||||||
|
"denoise_mask": "denoise_mask_function",
|
||||||
|
}
|
||||||
|
key = mode_keys.get(patch_mode)
|
||||||
|
if key is None:
|
||||||
|
raise ValueError(f"Bad mode: {patch_mode}")
|
||||||
|
if existing_patch_mode != "normal":
|
||||||
|
old_handlers = model.model_options.pop(key, None)
|
||||||
|
if old_handlers is None:
|
||||||
|
old_handlers = ()
|
||||||
|
else:
|
||||||
|
old_handlers = ()
|
||||||
|
|
||||||
|
def get_refs(*args, **kwargs) -> FilterRefs:
|
||||||
|
old_results = []
|
||||||
|
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
|
||||||
|
argdict = args[0]
|
||||||
|
elif patch_mode == "denoise_mask":
|
||||||
|
argdict = {
|
||||||
|
"sigma": args[0],
|
||||||
|
"denoise_mask": args[1].clone(),
|
||||||
|
"sigmas": kwargs["extra_options"]["sigmas"].clone(),
|
||||||
|
}
|
||||||
|
elif patch_mode == "apply_model":
|
||||||
|
argdict = args[1] | {"apply_function": args[0]}
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||||
|
if old_handlers:
|
||||||
|
ridx = 0 if existing_patch_mode == "extract_sequence" else -1
|
||||||
|
if patch_mode in {"cfg", "denoise_mask", "apply_model"}:
|
||||||
|
old_results = (old_handlers[0](*args, **kwargs),)
|
||||||
|
elif patch_mode == "pre_cfg":
|
||||||
|
old_results = [argdict["conds_out"]]
|
||||||
|
for hf in old_handlers:
|
||||||
|
result = hf(argdict | {"conds_out": old_results[ridx]}).copy()
|
||||||
|
if len(old_results) > 1 and existing_patch_mode != "extract":
|
||||||
|
old_results[1] = result
|
||||||
|
else:
|
||||||
|
old_results.append(result)
|
||||||
|
old_results = old_results[1:]
|
||||||
|
elif patch_mode == "post_cfg":
|
||||||
|
old_results = [argdict["denoised"].clone()]
|
||||||
|
for hf in old_handlers:
|
||||||
|
result = hf(argdict | {"denoised": old_results[ridx].clone()})
|
||||||
|
if len(old_results) > 1 and existing_patch_mode != "extract":
|
||||||
|
old_results[1] = result
|
||||||
|
else:
|
||||||
|
old_results.append(result)
|
||||||
|
old_results = old_results[1:]
|
||||||
|
kvs = {
|
||||||
|
"sigma": argdict["sigma"].clone(),
|
||||||
|
"sigma_float": argdict["sigma"].max().item(),
|
||||||
|
"old_results": tuple(old_results),
|
||||||
|
}
|
||||||
|
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
|
||||||
|
kvs |= {
|
||||||
|
"x": argdict["input"].clone(),
|
||||||
|
"cfg_scale": argdict["cond_scale"],
|
||||||
|
}
|
||||||
|
if patch_mode in {"post_cfg", "cfg"}:
|
||||||
|
kvs["cond"] = argdict["cond_denoised"].clone()
|
||||||
|
uncond = argdict.get("uncond_denoised", None)
|
||||||
|
kvs["uncond"] = uncond if uncond is None else uncond.clone()
|
||||||
|
if patch_mode == "post_cfg":
|
||||||
|
kvs["denoised"] = argdict["denoised"].clone()
|
||||||
|
else:
|
||||||
|
conds_out = argdict["conds_out"]
|
||||||
|
kvs["cond"] = conds_out[0].clone()
|
||||||
|
kvs["uncond"] = (
|
||||||
|
conds_out[1].clone()
|
||||||
|
if len(conds_out) > 1 and conds_out[1] is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
kvs["conds_out"] = list(conds_out)
|
||||||
|
elif patch_mode == "denoise_mask":
|
||||||
|
kvs |= {
|
||||||
|
"sigmas": argdict["sigmas"],
|
||||||
|
"denoise_mask": argdict["denoise_mask"],
|
||||||
|
}
|
||||||
|
elif patch_mode == "apply_model":
|
||||||
|
kvs |= {
|
||||||
|
"x": argdict["input"].clone(),
|
||||||
|
"cond_or_uncond": argdict["cond_or_uncond"].clone(),
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||||
|
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
|
||||||
|
latent_in = kvs["x"]
|
||||||
|
else:
|
||||||
|
latent_in = kvs["denoise_mask"]
|
||||||
|
kvs |= {k: v.to(latent_in) for k, v in latent_refs.items()}
|
||||||
|
return FilterRefs(kvs=kvs)
|
||||||
|
|
||||||
|
def model_patch(*args, **kwargs):
|
||||||
|
refs = get_refs(*args, **kwargs)
|
||||||
|
if patch_mode == "apply_model":
|
||||||
|
|
||||||
|
def fallback_apply_model():
|
||||||
|
old_results = refs.kvs["old_results"]
|
||||||
|
if old_results:
|
||||||
|
return old_results[-1]
|
||||||
|
return args[0](
|
||||||
|
args[1]["input"], args[1]["timestep"], **args[1]["c"]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
fallback_apply_model = None
|
||||||
|
if not ocs_filter.check_applies(refs):
|
||||||
|
old_results = refs.kvs["old_results"]
|
||||||
|
if old_results:
|
||||||
|
return old_results[-1]
|
||||||
|
if patch_mode == "pre_cfg":
|
||||||
|
return args[0]["conds_out"]
|
||||||
|
if patch_mode == "post_cfg":
|
||||||
|
return args[0]["denoised"]
|
||||||
|
if patch_mode == "cfg":
|
||||||
|
return args[0]["cond"]
|
||||||
|
if patch_mode == "denoise_mask":
|
||||||
|
return args[1]
|
||||||
|
if patch_mode == "apply_model":
|
||||||
|
return fallback_apply_model()
|
||||||
|
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||||
|
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
|
||||||
|
latent_in = refs.kvs["x"]
|
||||||
|
else:
|
||||||
|
latent_in = refs.kvs["denoise_mask"]
|
||||||
|
result = ocs_filter.apply(latent_in, refs=refs)
|
||||||
|
if patch_mode == "apply_model" and result is None:
|
||||||
|
return fallback_apply_model()
|
||||||
|
if patch_mode == "pre_cfg":
|
||||||
|
return list(result)
|
||||||
|
return result
|
||||||
|
|
||||||
|
if patch_mode == "pre_cfg":
|
||||||
|
model.set_model_sampler_pre_cfg_function(model_patch)
|
||||||
|
elif patch_mode == "post_cfg":
|
||||||
|
model.set_model_sampler_post_cfg_function(model_patch)
|
||||||
|
elif patch_mode == "cfg":
|
||||||
|
model.set_model_sampler_cfg_function(model_patch)
|
||||||
|
elif patch_mode == "denoise_mask":
|
||||||
|
model.set_model_denoise_mask_function(model_patch)
|
||||||
|
elif patch_mode == "apply_model":
|
||||||
|
model.set_model_unet_function_wrapper(model_patch)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||||
|
return (model,)
|
||||||
|
|
||||||
|
|
||||||
__all__ = (
|
__all__ = (
|
||||||
"SamplerNode",
|
"SamplerNode",
|
||||||
"GroupNode",
|
"GroupNode",
|
||||||
|
|||||||
+84
-19
@@ -2,12 +2,67 @@ import gc
|
|||||||
import math
|
import math
|
||||||
import random
|
import random
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import scipy
|
import scipy
|
||||||
import torch
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from .filtering import Filter, make_filter
|
from .filtering import Filter, make_filter
|
||||||
from .utils import scale_noise, fallback
|
from .utils import fallback, scale_noise
|
||||||
|
|
||||||
|
try:
|
||||||
|
from .triton_lsa import (
|
||||||
|
assignments_to_indices,
|
||||||
|
batch_linear_assignment,
|
||||||
|
batch_linear_assignment_shuffled,
|
||||||
|
)
|
||||||
|
|
||||||
|
HAVE_TRITON = True
|
||||||
|
except Exception:
|
||||||
|
HAVE_TRITON = False
|
||||||
|
|
||||||
|
|
||||||
|
def linear_sum_assignment(
|
||||||
|
cost: torch.Tensor,
|
||||||
|
*,
|
||||||
|
maximize: bool = False,
|
||||||
|
use_triton: bool = False,
|
||||||
|
split_batch: int = 0,
|
||||||
|
**kwargs: dict,
|
||||||
|
) -> tuple[np.ndarray, np.ndarray] | tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
if not use_triton or not HAVE_TRITON or not cost.is_cuda:
|
||||||
|
cost = cost.half().cpu()
|
||||||
|
return scipy.optimize.linear_sum_assignment(cost, maximize=maximize)
|
||||||
|
ndim = cost.ndim
|
||||||
|
orig_shape = cost.shape
|
||||||
|
if ndim == 2:
|
||||||
|
do_split = split_batch > 1 and all(
|
||||||
|
(sz / split_batch).is_integer() for sz in orig_shape
|
||||||
|
)
|
||||||
|
if do_split:
|
||||||
|
cost = cost.reshape(
|
||||||
|
split_batch, orig_shape[0] // split_batch, orig_shape[1] // split_batch
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
cost = cost.unsqueeze(0)
|
||||||
|
tqdm.write(
|
||||||
|
f"TRITON LAP: maximize={maximize}, orig cost shape={orig_shape}, cost shape={cost.shape}, cost dtype={cost.dtype}",
|
||||||
|
)
|
||||||
|
if not cost.is_contiguous():
|
||||||
|
cost = cost.contiguous()
|
||||||
|
fun = (
|
||||||
|
batch_linear_assignment
|
||||||
|
if "generator" not in kwargs
|
||||||
|
else batch_linear_assignment_shuffled
|
||||||
|
)
|
||||||
|
assignments = fun(cost, maximize=maximize, **kwargs)
|
||||||
|
row_ind, col_ind = assignments_to_indices(assignments)
|
||||||
|
if ndim == 2:
|
||||||
|
row_ind = row_ind.reshape(-1, row_ind.shape[-1])
|
||||||
|
col_ind = col_ind.reshape(-1, col_ind.shape[-1])
|
||||||
|
# row_ind, col_ind = row_ind.squeeze(0), col_ind.squeeze(0)
|
||||||
|
tqdm.write(f"Ran LAP kernel: {assignments.shape}, {row_ind.shape}, {col_ind.shape}")
|
||||||
|
return row_ind, col_ind
|
||||||
|
|
||||||
|
|
||||||
class ImmiscibleNoise(Filter):
|
class ImmiscibleNoise(Filter):
|
||||||
@@ -19,6 +74,12 @@ class ImmiscibleNoise(Filter):
|
|||||||
"maximize": False,
|
"maximize": False,
|
||||||
"distance_scale": 0.0,
|
"distance_scale": 0.0,
|
||||||
"distance_scale_ref": None,
|
"distance_scale_ref": None,
|
||||||
|
"abs_mode": False,
|
||||||
|
"abs_distance_mode": False,
|
||||||
|
"use_triton": False,
|
||||||
|
# Only honored in Triton mode.
|
||||||
|
"split_batch": 0,
|
||||||
|
"generator": None,
|
||||||
}
|
}
|
||||||
|
|
||||||
def __call__(self, noise_sampler, x_ref, *, refs=None):
|
def __call__(self, noise_sampler, x_ref, *, refs=None):
|
||||||
@@ -83,7 +144,11 @@ class ImmiscibleNoise(Filter):
|
|||||||
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
|
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
|
||||||
# Minimize latent-noise pairs over a batch
|
# Minimize latent-noise pairs over a batch
|
||||||
batch = latent.shape[0]
|
batch = latent.shape[0]
|
||||||
|
out_latent = fallback(out_latent, latent)
|
||||||
ref_latent = ref_latent.detach().clone()
|
ref_latent = ref_latent.detach().clone()
|
||||||
|
if self.abs_mode:
|
||||||
|
ref_latent = ref_latent.abs()
|
||||||
|
latent = latent.abs()
|
||||||
if self.distance_scale == 0:
|
if self.distance_scale == 0:
|
||||||
ref_latent_expanded = ref_latent.unsqueeze(1).expand(
|
ref_latent_expanded = ref_latent.unsqueeze(1).expand(
|
||||||
-1, batch, *ref_latent.shape[1:]
|
-1, batch, *ref_latent.shape[1:]
|
||||||
@@ -92,30 +157,30 @@ class ImmiscibleNoise(Filter):
|
|||||||
ref_latent.shape[0], *latent.shape
|
ref_latent.shape[0], *latent.shape
|
||||||
)
|
)
|
||||||
dist = (ref_latent_expanded - latent_expanded) ** 2
|
dist = (ref_latent_expanded - latent_expanded) ** 2
|
||||||
|
if self.abs_distance_mode:
|
||||||
|
dist = dist.abs_()
|
||||||
del ref_latent_expanded, latent_expanded
|
del ref_latent_expanded, latent_expanded
|
||||||
dist = dist.mean(tuple(range(2, dist.dim())))
|
dist = dist.mean(tuple(range(2, dist.ndim)))
|
||||||
else:
|
else:
|
||||||
dist = torch.linalg.vector_norm(
|
distance_scale_ref = fallback(self.distance_scale_ref, self.distance_scale)
|
||||||
fallback(self.distance_scale_ref, self.distance_scale)
|
dist = distance_scale_ref * ref_latent.flatten(start_dim=1).unsqueeze(
|
||||||
* ref_latent.flatten(start_dim=1).unsqueeze(1)
|
1
|
||||||
- self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0),
|
) - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0)
|
||||||
dim=2,
|
if self.abs_distance_mode:
|
||||||
)
|
dist = dist.abs_()
|
||||||
dist = dist.half()
|
dist = torch.linalg.vector_norm(dist, dim=2)
|
||||||
try:
|
try:
|
||||||
assign_mat = scipy.optimize.linear_sum_assignment(
|
assign_mat = linear_sum_assignment(
|
||||||
dist.cpu(), maximize=self.maximize
|
dist,
|
||||||
|
maximize=self.maximize,
|
||||||
|
use_triton=self.use_triton,
|
||||||
|
split_batch=self.split_batch,
|
||||||
|
generator=self.generator,
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}")
|
tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}")
|
||||||
return (
|
return None if return_idxs else out_latent[: ref_latent.shape[0]]
|
||||||
None
|
return assign_mat if return_idxs else out_latent[assign_mat[1]]
|
||||||
if return_idxs
|
|
||||||
else fallback(out_latent, latent)[: ref_latent.shape[0]]
|
|
||||||
)
|
|
||||||
return (
|
|
||||||
assign_mat if return_idxs else fallback(out_latent, latent)[assign_mat[1]]
|
|
||||||
)
|
|
||||||
|
|
||||||
def immiscible_simple(
|
def immiscible_simple(
|
||||||
self,
|
self,
|
||||||
|
|||||||
+112
-10
@@ -1,8 +1,57 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import NamedTuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
# from tqdm import tqdm
|
||||||
|
|
||||||
|
|
||||||
|
class RestartScaleFactors(NamedTuple):
|
||||||
|
latent_scale: float
|
||||||
|
noise_scale: float
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def build(
|
||||||
|
cls,
|
||||||
|
sigma_from: float | torch.Tensor,
|
||||||
|
sigma_to: float | torch.Tensor,
|
||||||
|
*,
|
||||||
|
is_flow: bool,
|
||||||
|
) -> RestartScaleFactors:
|
||||||
|
if isinstance(sigma_from, torch.Tensor):
|
||||||
|
sigma_from = sigma_from.max().item()
|
||||||
|
if isinstance(sigma_to, torch.Tensor):
|
||||||
|
sigma_to = sigma_to.max().item()
|
||||||
|
if not is_flow:
|
||||||
|
return cls(
|
||||||
|
1.0,
|
||||||
|
max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5,
|
||||||
|
)
|
||||||
|
alpha_from = 1.0 - sigma_from
|
||||||
|
alpha_to = 1.0 - sigma_to
|
||||||
|
if alpha_to <= 0:
|
||||||
|
latent_scale = 0.0
|
||||||
|
noise_scale = sigma_to
|
||||||
|
else:
|
||||||
|
latent_scale = alpha_to / alpha_from
|
||||||
|
noise_scale = (
|
||||||
|
max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5
|
||||||
|
)
|
||||||
|
return cls(latent_scale, noise_scale)
|
||||||
|
|
||||||
|
|
||||||
class Restart:
|
class Restart:
|
||||||
def __init__(self, *, s_noise=1.0, custom_noise=None, immiscible=False):
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
s_noise=1.0,
|
||||||
|
custom_noise=None,
|
||||||
|
immiscible=False,
|
||||||
|
normalized=True,
|
||||||
|
normalize_dims: tuple[int, ...] | None = None,
|
||||||
|
is_flow=False,
|
||||||
|
):
|
||||||
from .noise import ImmiscibleNoise
|
from .noise import ImmiscibleNoise
|
||||||
|
|
||||||
self.s_noise = s_noise
|
self.s_noise = s_noise
|
||||||
@@ -10,6 +59,9 @@ class Restart:
|
|||||||
immiscible = ImmiscibleNoise(**immiscible)
|
immiscible = ImmiscibleNoise(**immiscible)
|
||||||
self.immiscible = immiscible
|
self.immiscible = immiscible
|
||||||
self.custom_noise = custom_noise
|
self.custom_noise = custom_noise
|
||||||
|
self.normalized = normalized
|
||||||
|
self.normalize_dims = normalize_dims
|
||||||
|
self.is_flow = is_flow
|
||||||
|
|
||||||
def get_noise_sampler(self, nsc):
|
def get_noise_sampler(self, nsc):
|
||||||
return nsc.make_caching_noise_sampler(
|
return nsc.make_caching_noise_sampler(
|
||||||
@@ -30,23 +82,64 @@ class Restart:
|
|||||||
last_sigma = sigma
|
last_sigma = sigma
|
||||||
return sigmas
|
return sigmas
|
||||||
|
|
||||||
def split_sigmas(self, sigmas):
|
def split_sigmas(self, sigmas: torch.Tensor):
|
||||||
prev_seg = None
|
prev_seg = None
|
||||||
while len(sigmas) > 1:
|
while len(sigmas) > 1:
|
||||||
seg = self.get_segment(sigmas)
|
seg = self.get_segment(sigmas)
|
||||||
sigmas = sigmas[len(seg) :]
|
sigmas = sigmas[len(seg) :]
|
||||||
if prev_seg is not None and seg[0] > prev_seg[-1]:
|
if prev_seg is not None and seg[0] > prev_seg[-1]:
|
||||||
noise_scale = self.get_noise_scale(prev_seg[-1], seg[0])
|
scale_factors = RestartScaleFactors.build(
|
||||||
|
sigma_from=prev_seg[-1], sigma_to=seg[0], is_flow=self.is_flow
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
noise_scale = 0.0
|
scale_factors = None
|
||||||
prev_seg = seg
|
prev_seg = seg
|
||||||
yield (noise_scale, seg)
|
yield (scale_factors, seg)
|
||||||
|
|
||||||
def get_noise_scale(self, s_min, s_max):
|
def get_noise_scale(
|
||||||
|
self, s_min: float | torch.Tensor, s_max: float | torch.Tensor
|
||||||
|
) -> float:
|
||||||
result = (s_max**2 - s_min**2) ** 0.5
|
result = (s_max**2 - s_min**2) ** 0.5
|
||||||
if isinstance(result, torch.Tensor):
|
if isinstance(result, torch.Tensor):
|
||||||
result = result.item()
|
return result.item()
|
||||||
return result * self.s_noise
|
return result
|
||||||
|
|
||||||
|
def add_noise(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
sigma_from: float,
|
||||||
|
sigma_to: float,
|
||||||
|
*,
|
||||||
|
nsc,
|
||||||
|
refs,
|
||||||
|
scale_factors: RestartScaleFactors | None = None,
|
||||||
|
in_place: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if self.is_flow:
|
||||||
|
sigma_from = min(1.0, max(0.0, sigma_from))
|
||||||
|
sigma_to = min(1.0, max(0.0, sigma_to))
|
||||||
|
if sigma_from >= sigma_to:
|
||||||
|
raise ValueError(
|
||||||
|
f"sigma_from ({sigma_from:.4f}) must be less than sigma_to ({sigma_to:.4f})"
|
||||||
|
)
|
||||||
|
scale_factors = scale_factors or RestartScaleFactors.build(
|
||||||
|
sigma_from, sigma_to, is_flow=self.is_flow
|
||||||
|
)
|
||||||
|
ns = self.get_noise_sampler(nsc)
|
||||||
|
sigma_empty = nsc.min_sigma * 0
|
||||||
|
noise = nsc.scale_noise(
|
||||||
|
ns(sigma_empty + sigma_from, sigma_empty + sigma_to, refs=refs),
|
||||||
|
normalized=self.normalized,
|
||||||
|
normalize_dims=self.normalize_dims,
|
||||||
|
)
|
||||||
|
noise *= scale_factors.noise_scale * self.s_noise
|
||||||
|
if scale_factors.latent_scale != 1.0:
|
||||||
|
x = (
|
||||||
|
x.mul_(scale_factors.latent_scale)
|
||||||
|
if in_place
|
||||||
|
else scale_factors.latent_scale * x
|
||||||
|
)
|
||||||
|
return noise.add_(x)
|
||||||
|
|
||||||
def __repr__(self):
|
def __repr__(self):
|
||||||
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
|
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
|
||||||
@@ -77,18 +170,27 @@ class Restart:
|
|||||||
raise ValueError("Schedule jump index out of range")
|
raise ValueError("Schedule jump index out of range")
|
||||||
sched_idx = item
|
sched_idx = item
|
||||||
continue
|
continue
|
||||||
|
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
|
||||||
|
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
|
||||||
if sig_idx >= siglen or sig_idx < 0:
|
if sig_idx >= siglen or sig_idx < 0:
|
||||||
break
|
break
|
||||||
interval, jump = item
|
interval, jump = item
|
||||||
chunk = siglist[sig_idx : sig_idx + interval + 1]
|
chunk = siglist[sig_idx : sig_idx + interval + 1]
|
||||||
|
if sched_frac != 0:
|
||||||
|
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
|
||||||
# print(f"{out} + {chunk}")
|
# print(f"{out} + {chunk}")
|
||||||
out += chunk
|
out += chunk
|
||||||
sig_idx += interval + jump
|
|
||||||
if jump >= 0:
|
if jump >= 0:
|
||||||
sig_idx += 1
|
sig_idx += 1
|
||||||
|
sig_idx += interval + jump
|
||||||
sched_idx += 1
|
sched_idx += 1
|
||||||
|
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
|
||||||
|
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
|
||||||
if sig_idx < siglen and sig_idx >= 0:
|
if sig_idx < siglen and sig_idx >= 0:
|
||||||
out += siglist[sig_idx:]
|
chunk = siglist[sig_idx:]
|
||||||
|
if sched_frac != 0:
|
||||||
|
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
|
||||||
|
out += chunk
|
||||||
if out[-1] > siglist[-1]:
|
if out[-1] > siglist[-1]:
|
||||||
out.append(siglist[-1])
|
out.append(siglist[-1])
|
||||||
return torch.tensor(out).to(sigmas)
|
return torch.tensor(out).to(sigmas)
|
||||||
|
|||||||
+31
-21
@@ -1,13 +1,12 @@
|
|||||||
import torch
|
import torch
|
||||||
from tqdm.auto import trange
|
from tqdm.auto import trange
|
||||||
|
|
||||||
|
|
||||||
from .filtering import FILTER_HANDLERS, FilterRefs
|
from .filtering import FILTER_HANDLERS, FilterRefs
|
||||||
from .model import OCSModel
|
from .model import OCSModel
|
||||||
from .noise import NoiseSamplerCache
|
from .noise import NoiseSamplerCache
|
||||||
from .substep_sampling import SamplerState
|
|
||||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
|
||||||
from .restart import Restart
|
from .restart import Restart
|
||||||
|
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||||
|
from .substep_sampling import SamplerState
|
||||||
|
|
||||||
|
|
||||||
def find_merge_sampler(merge_samplers, ss) -> object | None:
|
def find_merge_sampler(merge_samplers, ss) -> object | None:
|
||||||
@@ -48,11 +47,6 @@ def composable_sampler(
|
|||||||
restart_custom_noise = copts.get("restart_custom_noise")
|
restart_custom_noise = copts.get("restart_custom_noise")
|
||||||
if isinstance(restart_custom_noise, str):
|
if isinstance(restart_custom_noise, str):
|
||||||
restart_custom_noise = copts.get(f"restart_custom_noise_{restart_custom_noise}")
|
restart_custom_noise = copts.get(f"restart_custom_noise_{restart_custom_noise}")
|
||||||
restart = Restart(
|
|
||||||
s_noise=restart_params.get("s_noise", 1.0),
|
|
||||||
custom_noise=restart_custom_noise,
|
|
||||||
immiscible=restart_params.get("immiscible", False),
|
|
||||||
)
|
|
||||||
|
|
||||||
ss = SamplerState(
|
ss = SamplerState(
|
||||||
OCSModel(
|
OCSModel(
|
||||||
@@ -72,6 +66,16 @@ def composable_sampler(
|
|||||||
reta=copts.get("reta", 1.0),
|
reta=copts.get("reta", 1.0),
|
||||||
disable_status=disable,
|
disable_status=disable,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
restart = Restart(
|
||||||
|
s_noise=restart_params.get("s_noise", 1.0),
|
||||||
|
custom_noise=restart_custom_noise,
|
||||||
|
immiscible=restart_params.get("immiscible", False),
|
||||||
|
normalized=restart_params.get("normalized", True),
|
||||||
|
normalize_dims=restart_params.get("normalize_dims"),
|
||||||
|
is_flow=ss.model.is_rectified_flow,
|
||||||
|
)
|
||||||
|
|
||||||
groups = copts["_groups"]
|
groups = copts["_groups"]
|
||||||
merge_samplers = tuple(
|
merge_samplers = tuple(
|
||||||
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
|
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
|
||||||
@@ -85,17 +89,17 @@ def composable_sampler(
|
|||||||
)
|
)
|
||||||
ss.noise = nsc
|
ss.noise = nsc
|
||||||
sigma_chunks = (
|
sigma_chunks = (
|
||||||
tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((0.0, sigmas),)
|
tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((None, sigmas),)
|
||||||
)
|
)
|
||||||
step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks)
|
step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks)
|
||||||
ss.total_steps = step_count
|
ss.total_steps = step_count
|
||||||
step = 0
|
step = 0
|
||||||
with trange(step_count, disable=ss.disable_status) as pbar:
|
with trange(step_count, disable=ss.disable_status) as pbar:
|
||||||
for noise_scale, chunk_sigmas in sigma_chunks:
|
for chunk_idx, (scale_factors, chunk_sigmas) in enumerate(sigma_chunks):
|
||||||
if step != 0 and noise_scale != 0:
|
if step != 0 and scale_factors is not None:
|
||||||
prev_refs = FilterRefs({
|
prev_refs = FilterRefs(
|
||||||
f"pre_restart_{k}": v for k, v in ss.refs.items()
|
{f"pre_restart_{k}": v for k, v in ss.refs.items()}
|
||||||
})
|
)
|
||||||
ss.sigmas = chunk_sigmas
|
ss.sigmas = chunk_sigmas
|
||||||
ss.update(0, step=step, substep=0)
|
ss.update(0, step=step, substep=0)
|
||||||
if step != 0:
|
if step != 0:
|
||||||
@@ -104,14 +108,20 @@ def composable_sampler(
|
|||||||
ss.hist.reset()
|
ss.hist.reset()
|
||||||
for ms in merge_samplers:
|
for ms in merge_samplers:
|
||||||
ms.reset()
|
ms.reset()
|
||||||
nsc.min_sigma, nsc.max_sigma = chunk_sigmas[-1], chunk_sigmas[0]
|
nsc.min_sigma, nsc.max_sigma = (
|
||||||
if step != 0 and noise_scale != 0:
|
chunk_sigmas[-1].clone(),
|
||||||
restart_ns = restart.get_noise_sampler(nsc)
|
chunk_sigmas[0].clone(),
|
||||||
x += nsc.scale_noise(
|
)
|
||||||
restart_ns(nsc.min_sigma, nsc.max_sigma, refs=prev_refs | ss.refs),
|
if step != 0 and scale_factors is not None:
|
||||||
noise_scale,
|
x = restart.add_noise(
|
||||||
|
x,
|
||||||
|
sigma_from=sigma_chunks[chunk_idx - 1][1][-1].item(),
|
||||||
|
sigma_to=chunk_sigmas[0].item(),
|
||||||
|
scale_factors=scale_factors,
|
||||||
|
nsc=nsc,
|
||||||
|
refs=prev_refs | ss.refs,
|
||||||
|
in_place=True,
|
||||||
)
|
)
|
||||||
del restart_ns
|
|
||||||
del prev_refs
|
del prev_refs
|
||||||
for idx in range(len(chunk_sigmas) - 1):
|
for idx in range(len(chunk_sigmas) - 1):
|
||||||
if idx > 0:
|
if idx > 0:
|
||||||
|
|||||||
@@ -1,20 +1,19 @@
|
|||||||
|
import inspect
|
||||||
|
import math
|
||||||
import typing
|
import typing
|
||||||
|
|
||||||
import inspect
|
import comfy
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
import comfy
|
|
||||||
|
|
||||||
from .. import filtering
|
|
||||||
from .. import expression as expr
|
from .. import expression as expr
|
||||||
|
from .. import filtering
|
||||||
from ..utils import fallback
|
from ..utils import fallback
|
||||||
from .base import (
|
from .base import (
|
||||||
StepSamplerContext,
|
|
||||||
SingleStepSampler,
|
SingleStepSampler,
|
||||||
|
StepSamplerContext,
|
||||||
registry,
|
registry,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import pytorch_wavelets as ptwav
|
import pytorch_wavelets as ptwav
|
||||||
|
|
||||||
@@ -531,7 +530,10 @@ class WeoonStep(SingleStepSampler):
|
|||||||
)
|
)
|
||||||
denoised_new = self.wavelet_inverse(coeffs_out)
|
denoised_new = self.wavelet_inverse(coeffs_out)
|
||||||
if denoised_new.shape != x.shape:
|
if denoised_new.shape != x.shape:
|
||||||
denoised_new = denoised_new.reshape(*x.shape)
|
bi_elements = math.prod(x.shape[1:])
|
||||||
|
denoised_new = denoised_new.reshape(x.shape[0], -1)[
|
||||||
|
:, :bi_elements
|
||||||
|
].reshape(*x.shape)
|
||||||
x = self.blend(denoised_new, x, ratio)
|
x = self.blend(denoised_new, x, ratio)
|
||||||
yield from self.result(x, sigma_up)
|
yield from self.result(x, sigma_up)
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
import torch
|
|
||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
|
import torch
|
||||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
from .. import filtering
|
||||||
from .base import (
|
from .base import (
|
||||||
SingleStepSampler,
|
|
||||||
DPMPPStepMixin,
|
DPMPPStepMixin,
|
||||||
HistorySingleStepSampler,
|
HistorySingleStepSampler,
|
||||||
ReversibleSingleStepSampler,
|
ReversibleSingleStepSampler,
|
||||||
|
SingleStepSampler,
|
||||||
registry,
|
registry,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -524,6 +525,147 @@ class RESMultistepStep(HistorySingleStepSampler, DPMPPStepMixin):
|
|||||||
yield from self.result(result, sigma_up, sigma_down=sigma_down)
|
yield from self.result(result, sigma_up, sigma_down=sigma_down)
|
||||||
|
|
||||||
|
|
||||||
|
# SEEDS-2 - Stochastic Explicit Exponential Derivative-free Solvers (VP Data Prediction) stage 2.
|
||||||
|
# arXiv: https://arxiv.org/abs/2305.14267 (NeurIPS 2023)
|
||||||
|
# Implementation referenced from ComfyUI.
|
||||||
|
class Seeds2Step(SingleStepSampler, DPMPPStepMixin):
|
||||||
|
name = "seeds_2"
|
||||||
|
self_noise = 3
|
||||||
|
model_calls = 1
|
||||||
|
allow_alt_cfgpp = False
|
||||||
|
uses_alt_noise = True
|
||||||
|
|
||||||
|
def __init__(self, *args, r=0.5, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.r = r
|
||||||
|
s2_options = self.options.get("seeds_2", {})
|
||||||
|
sigma_blend_mode = s2_options.get("sigma_blend_mode", "lerp").strip()
|
||||||
|
self.sigma_blend_function = (
|
||||||
|
filtering.BLENDING_MODES[sigma_blend_mode]
|
||||||
|
if sigma_blend_mode != "lerp"
|
||||||
|
else torch.lerp
|
||||||
|
)
|
||||||
|
denoised_blend_mode = s2_options.get("denoised_blend_mode", "lerp").strip()
|
||||||
|
self.denoised_blend_function = (
|
||||||
|
filtering.BLENDING_MODES[denoised_blend_mode]
|
||||||
|
if denoised_blend_mode != "lerp"
|
||||||
|
else torch.lerp
|
||||||
|
)
|
||||||
|
self.disable_stage2_eta = bool(s2_options.get("disable_stage2_eta", False))
|
||||||
|
stage2_stage1_noise_blend_mode = s2_options.get(
|
||||||
|
"stage2_stage1_noise_blend_mode", "lerp"
|
||||||
|
).strip()
|
||||||
|
self.stage2_stage1_noise_blend_function = (
|
||||||
|
filtering.BLENDING_MODES[stage2_stage1_noise_blend_mode]
|
||||||
|
if stage2_stage1_noise_blend_mode != "lerp"
|
||||||
|
else torch.lerp
|
||||||
|
)
|
||||||
|
self.stage2_stage1_noise_ratio = s2_options.get(
|
||||||
|
"stage2_stage1_noise_ratio", 1.0
|
||||||
|
)
|
||||||
|
self.stage1_s_noise = s2_options.get("stage1_s_noise", 1.0)
|
||||||
|
self.stage2_s_noise = s2_options.get("stage2_s_noise", 1.0)
|
||||||
|
self.stage2_sigma_scale = s2_options.get("stage2_sigma_scale", 1.0)
|
||||||
|
|
||||||
|
def step(self, x: torch.Tensor):
|
||||||
|
ss = self.ss
|
||||||
|
sigma = ss.sigma.to(dtype=torch.float64)
|
||||||
|
sigma_next = ss.sigma_next.to(dtype=torch.float64)
|
||||||
|
denoised = ss.denoised
|
||||||
|
|
||||||
|
t_one = ss.sigma * 0 + 1.0
|
||||||
|
|
||||||
|
r, eta = self.r, self.get_dyn_eta()
|
||||||
|
fac = 1 / (2 * r)
|
||||||
|
lambda_s = ss.sigma_to_half_log_snr(sigma=sigma)
|
||||||
|
lambda_t = ss.sigma_to_half_log_snr(sigma=sigma_next)
|
||||||
|
h = lambda_t - lambda_s
|
||||||
|
h_eta = h * (eta + 1.0)
|
||||||
|
lambda_s_1 = self.sigma_blend_function(
|
||||||
|
lambda_s.unsqueeze(0), lambda_t.unsqueeze(0), r
|
||||||
|
).squeeze(0)
|
||||||
|
sigma_s_1 = ss.half_log_snr_to_sigma(lambda_s_1)
|
||||||
|
|
||||||
|
alpha_s_1 = sigma_s_1 * lambda_s_1.exp()
|
||||||
|
alpha_t = sigma_next * lambda_t.exp()
|
||||||
|
|
||||||
|
s1_x_mult = sigma_s_1 / sigma * (-r * h * eta).exp()
|
||||||
|
s1_denoised_mult = alpha_s_1 * (-r * h_eta).expm1()
|
||||||
|
x_2 = (
|
||||||
|
s1_x_mult.to(dtype=x.dtype) * x
|
||||||
|
- s1_denoised_mult.to(dtype=x.dtype) * denoised
|
||||||
|
)
|
||||||
|
if eta != 0:
|
||||||
|
s1_noise_mult = (-2 * r * h * eta).expm1().neg().sqrt()
|
||||||
|
sde_noise1 = yield from self.result(
|
||||||
|
x_2 * 0,
|
||||||
|
s1_noise_mult.to(dtype=x.dtype),
|
||||||
|
sigma=ss.sigma,
|
||||||
|
sigma_next=sigma_s_1.to(dtype=x.dtype),
|
||||||
|
noise_sampler=self.alt_noise_sampler,
|
||||||
|
final=False,
|
||||||
|
)
|
||||||
|
x_2 += sde_noise1 * (sigma_s_1.to(dtype=x.dtype) * self.stage1_s_noise)
|
||||||
|
|
||||||
|
denoised_2 = self.call_model(
|
||||||
|
x_2, (sigma_s_1 * self.stage2_sigma_scale).to(dtype=x.dtype), call_index=1
|
||||||
|
).denoised
|
||||||
|
denoised_d = self.denoised_blend_function(denoised, denoised_2, fac)
|
||||||
|
|
||||||
|
if self.disable_stage2_eta:
|
||||||
|
eta = 0.0
|
||||||
|
h_eta = h
|
||||||
|
|
||||||
|
s2_x_mult = sigma_next / sigma * (-h * eta).exp()
|
||||||
|
s2_denoised_mult = alpha_t * h_eta.neg().expm1()
|
||||||
|
x_curr = s2_x_mult.to(dtype=x.dtype) * x
|
||||||
|
x_curr -= s2_denoised_mult.to(dtype=x.dtype) * denoised_d
|
||||||
|
|
||||||
|
if eta == 0:
|
||||||
|
return (yield from self.result(x_curr))
|
||||||
|
|
||||||
|
s2_s1_nr = self.stage2_stage1_noise_ratio
|
||||||
|
|
||||||
|
segment_factor = ((r - 1.0) * h * eta).to(dtype=x.dtype)
|
||||||
|
s2_noise_mult = (segment_factor * 2.0).expm1().neg() ** 0.5
|
||||||
|
sde_noise2_raw = yield from self.result(
|
||||||
|
x_curr * 0,
|
||||||
|
t_one,
|
||||||
|
sigma=sigma_s_1.to(dtype=x.dtype),
|
||||||
|
sigma_next=ss.sigma_next,
|
||||||
|
final=False,
|
||||||
|
)
|
||||||
|
sde_noise2 = sde_noise2_raw * s2_noise_mult.to(dtype=x.dtype)
|
||||||
|
|
||||||
|
if s2_s1_nr != 1.0:
|
||||||
|
# print(
|
||||||
|
# f"\n\nBLENDING: {s2_s1_nr:.4f}, {s1_noise_mult.item():.4f}, {s2_noise_mult.item():.4f}"
|
||||||
|
# )
|
||||||
|
sde_noise1 = self.stage2_stage1_noise_blend_function(
|
||||||
|
(
|
||||||
|
yield from self.result(
|
||||||
|
x_curr * 0,
|
||||||
|
t_one,
|
||||||
|
sigma=sigma_s_1.to(dtype=x.dtype),
|
||||||
|
sigma_next=ss.sigma_next,
|
||||||
|
final=False,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
* s1_noise_mult.to(dtype=x.dtype),
|
||||||
|
sde_noise1,
|
||||||
|
s2_s1_nr,
|
||||||
|
)
|
||||||
|
|
||||||
|
sde_noise1 *= segment_factor.exp()
|
||||||
|
sde_noise2 += sde_noise1
|
||||||
|
sde_noise2 *= ss.sigma_next * self.stage2_s_noise
|
||||||
|
x_curr += sde_noise2
|
||||||
|
|
||||||
|
yield from self.result(
|
||||||
|
x_curr, noise_scale=ss.sigma * 0, sigma_down=ss.sigma_next
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
registry.add(
|
registry.add(
|
||||||
DEISStep,
|
DEISStep,
|
||||||
DPMPP2MSDEStep,
|
DPMPP2MSDEStep,
|
||||||
@@ -538,4 +680,5 @@ registry.add(
|
|||||||
DPM2Step,
|
DPM2Step,
|
||||||
DPMPP2SStep,
|
DPMPP2SStep,
|
||||||
RESMultistepStep,
|
RESMultistepStep,
|
||||||
|
Seeds2Step,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ class PingPongStep(SingleStepSampler):
|
|||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
pingpong_options = self.options.pop("pingpong", {})
|
pingpong_options = self.options.pop("pingpong", {})
|
||||||
self.pingpong_start_step = pingpong_options.get("start_step", 0)
|
self.pingpong_start_step = pingpong_options.get("start_step", 0)
|
||||||
self.pingpong_end_step = pingpong_options.get("end_step", 0)
|
self.pingpong_end_step = pingpong_options.get("end_step", 9999)
|
||||||
|
|
||||||
def step(self, x):
|
def step(self, x):
|
||||||
ss = self.ss
|
ss = self.ss
|
||||||
|
|||||||
+16
-11
@@ -5,8 +5,7 @@ import tqdm
|
|||||||
|
|
||||||
from . import expression as expr
|
from . import expression as expr
|
||||||
from . import utils
|
from . import utils
|
||||||
|
from .filtering import FILTER_HANDLERS, FilterRefs, make_filter
|
||||||
from .filtering import make_filter, FilterRefs, FILTER_HANDLERS
|
|
||||||
from .noise import ImmiscibleNoise
|
from .noise import ImmiscibleNoise
|
||||||
from .restart import Restart
|
from .restart import Restart
|
||||||
from .step_samplers import STEP_SAMPLERS
|
from .step_samplers import STEP_SAMPLERS
|
||||||
@@ -451,6 +450,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
|||||||
s_noise=restart.get("s_noise", 1.0),
|
s_noise=restart.get("s_noise", 1.0),
|
||||||
custom_noise=restart_custom_noise,
|
custom_noise=restart_custom_noise,
|
||||||
immiscible=restart.get("immiscible", False),
|
immiscible=restart.get("immiscible", False),
|
||||||
|
is_flow=ss.model.is_rectified_flow,
|
||||||
)
|
)
|
||||||
|
|
||||||
def make_schedule(self, ss):
|
def make_schedule(self, ss):
|
||||||
@@ -505,10 +505,13 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
|||||||
if subss.idx >= max_idx:
|
if subss.idx >= max_idx:
|
||||||
break
|
break
|
||||||
if last_down is not None and last_down < ss.sigma_next:
|
if last_down is not None and last_down < ss.sigma_next:
|
||||||
restart_ns = self.restart.get_noise_sampler(ss.noise)
|
x = self.restart.add_noise(
|
||||||
x += ss.noise.scale_noise(
|
x,
|
||||||
restart_ns(last_down, ss.sigma_next, refs=ss.refs),
|
sigma_from=last_down.item(),
|
||||||
self.restart.get_noise_scale(last_down, ss.sigma_next),
|
sigma_to=ss.sigma_next.item(),
|
||||||
|
nsc=nsc,
|
||||||
|
refs=ss.refs,
|
||||||
|
in_place=True,
|
||||||
)
|
)
|
||||||
pbar.update(0)
|
pbar.update(0)
|
||||||
return x
|
return x
|
||||||
@@ -653,11 +656,13 @@ class PingpongMergeSubstepsSampler(MergeSubstepsSampler):
|
|||||||
sigma_next,
|
sigma_next,
|
||||||
immiscible=fallback(self.immiscible, ss.noise.immiscible),
|
immiscible=fallback(self.immiscible, ss.noise.immiscible),
|
||||||
)
|
)
|
||||||
noise_refs = ss.refs | FilterRefs({
|
noise_refs = ss.refs | FilterRefs(
|
||||||
"orig_x": orig_x,
|
{
|
||||||
"x": x,
|
"orig_x": orig_x,
|
||||||
"denoised": synth_denoised,
|
"x": x,
|
||||||
})
|
"denoised": synth_denoised,
|
||||||
|
}
|
||||||
|
)
|
||||||
noise = (
|
noise = (
|
||||||
noise_sampler(sigma, sigma_next, refs=noise_refs) * self.pingpong_s_noise
|
noise_sampler(sigma, sigma_next, refs=noise_refs) * self.pingpong_s_noise
|
||||||
)
|
)
|
||||||
|
|||||||
+92
-8
@@ -1,5 +1,6 @@
|
|||||||
import torch
|
from typing import NamedTuple
|
||||||
|
|
||||||
|
import torch
|
||||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||||
|
|
||||||
from .filtering import FilterRefs
|
from .filtering import FilterRefs
|
||||||
@@ -7,6 +8,13 @@ from .model import History
|
|||||||
from .utils import fallback
|
from .utils import fallback
|
||||||
|
|
||||||
|
|
||||||
|
class AncestralRatios(NamedTuple):
|
||||||
|
alpha_t: torch.Tensor
|
||||||
|
alpha_s: torch.Tensor
|
||||||
|
sigma_up: torch.Tensor
|
||||||
|
sigma_down: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
class Items:
|
class Items:
|
||||||
def __init__(self, items=None):
|
def __init__(self, items=None):
|
||||||
self.items = [] if items is None else items
|
self.items = [] if items is None else items
|
||||||
@@ -141,6 +149,10 @@ class SamplerState:
|
|||||||
self.substep = 0
|
self.substep = 0
|
||||||
self.total_steps = len(sigmas) - 1
|
self.total_steps = len(sigmas) - 1
|
||||||
self.cfg_scale_override = cfg_scale_override
|
self.cfg_scale_override = cfg_scale_override
|
||||||
|
self.is_flow = self.model.is_rectified_flow
|
||||||
|
self.offset_sigma = (
|
||||||
|
model.model_sampling.percent_to_sigma(1e-04) if self.is_flow else None
|
||||||
|
)
|
||||||
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
|
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -171,6 +183,23 @@ class SamplerState:
|
|||||||
def d(self):
|
def d(self):
|
||||||
return self.hcur.d
|
return self.hcur.d
|
||||||
|
|
||||||
|
# These two functions referenced from ComfyUI.
|
||||||
|
def sigma_to_half_log_snr(
|
||||||
|
self, *, sigma: torch.Tensor | None = None, idx: int | None = None
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if sigma is None and idx is None:
|
||||||
|
sigma = self.sigma
|
||||||
|
else:
|
||||||
|
sigma = sigma if sigma is not None else self.sigmas[idx]
|
||||||
|
if not self.is_flow:
|
||||||
|
return sigma.log().neg_()
|
||||||
|
if sigma.max() >= 1.0:
|
||||||
|
sigma = sigma * 0.0 + self.offset_sigma
|
||||||
|
return sigma.logit().neg_()
|
||||||
|
|
||||||
|
def half_log_snr_to_sigma(self, half_log_snr: torch.Tensor) -> torch.Tensor:
|
||||||
|
return (torch.sigmoid if self.is_flow else torch.exp)(half_log_snr.neg())
|
||||||
|
|
||||||
def update(self, idx=None, step=None, substep=None):
|
def update(self, idx=None, step=None, substep=None):
|
||||||
idx = self.idx if idx is None else idx
|
idx = self.idx if idx is None else idx
|
||||||
self.idx = idx
|
self.idx = idx
|
||||||
@@ -185,6 +214,59 @@ class SamplerState:
|
|||||||
self.substep = substep
|
self.substep = substep
|
||||||
self.refs = FilterRefs.from_ss(self)
|
self.refs = FilterRefs.from_ss(self)
|
||||||
|
|
||||||
|
def get_ancestral_step_ext(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
sigma: torch.Tensor | None = None,
|
||||||
|
sigma_next: torch.Tensor | None = None,
|
||||||
|
eta: float = 1.0,
|
||||||
|
retry_increment: int = 0,
|
||||||
|
):
|
||||||
|
sigma = fallback(sigma, self.sigma)
|
||||||
|
sigma_next = fallback(sigma_next, self.sigma_next)
|
||||||
|
sigma_empty = sigma_next * 0.0
|
||||||
|
|
||||||
|
def get_noeta_ratios():
|
||||||
|
return AncestralRatios(
|
||||||
|
alpha_t=sigma_empty + 1.0,
|
||||||
|
alpha_s=sigma_empty + 1.0,
|
||||||
|
sigma_up=sigma_empty.clone(),
|
||||||
|
sigma_down=sigma_next.clone(),
|
||||||
|
)
|
||||||
|
|
||||||
|
if eta <= 0 or sigma_next.max().item() <= 1e-08:
|
||||||
|
return get_noeta_ratios()
|
||||||
|
orig_dtype = sigma.dtype
|
||||||
|
sigma = sigma.to(dtype=torch.float64)
|
||||||
|
sigma_next = sigma_next.to(dtype=torch.float64)
|
||||||
|
alpha_s = sigma * self.sigma_to_half_log_snr(sigma=sigma).exp()
|
||||||
|
alpha_t = sigma_next * self.sigma_to_half_log_snr(sigma=sigma_next).exp()
|
||||||
|
adj_sigma = sigma / alpha_s
|
||||||
|
adj_sigma_next = sigma_next / alpha_t
|
||||||
|
sd = su = None
|
||||||
|
while eta > 0:
|
||||||
|
sd, su = (
|
||||||
|
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
|
||||||
|
for v in get_ancestral_step(adj_sigma, adj_sigma_next, eta=eta)
|
||||||
|
)
|
||||||
|
if sd > 0 and su > 0:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
sd = su = None
|
||||||
|
if retry_increment <= 0:
|
||||||
|
break
|
||||||
|
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
|
||||||
|
eta -= retry_increment
|
||||||
|
if sd is None or su is None:
|
||||||
|
return get_noeta_ratios()
|
||||||
|
sd = alpha_t * sd
|
||||||
|
return AncestralRatios(
|
||||||
|
alpha_t=alpha_t.to(dtype=orig_dtype),
|
||||||
|
alpha_s=alpha_s.to(dtype=orig_dtype),
|
||||||
|
sigma_up=su.to(dtype=orig_dtype),
|
||||||
|
sigma_down=sd.to(dtype=orig_dtype),
|
||||||
|
)
|
||||||
|
|
||||||
def get_ancestral_step(
|
def get_ancestral_step(
|
||||||
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
|
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
|
||||||
):
|
):
|
||||||
@@ -265,13 +347,15 @@ class SamplerState:
|
|||||||
preview = (hi.x - hi.denoised) * 0.1 + hi.denoised
|
preview = (hi.x - hi.denoised) * 0.1 + hi.denoised
|
||||||
else:
|
else:
|
||||||
preview = hi.denoised
|
preview = hi.denoised
|
||||||
return self.callback_({
|
return self.callback_(
|
||||||
"x": hi.x,
|
{
|
||||||
"i": self.step,
|
"x": hi.x,
|
||||||
"sigma": hi.sigma,
|
"i": self.step,
|
||||||
"sigma_hat": hi.sigma,
|
"sigma": hi.sigma,
|
||||||
"denoised": preview,
|
"sigma_hat": hi.sigma,
|
||||||
})
|
"denoised": preview,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
self.hist.reset()
|
self.hist.reset()
|
||||||
|
|||||||
@@ -0,0 +1,386 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
@triton.autotune(
|
||||||
|
configs=[
|
||||||
|
triton.Config({"num_warps": 4, "num_stages": 2}, num_warps=4, num_stages=2),
|
||||||
|
triton.Config({"num_warps": 8, "num_stages": 2}, num_warps=8, num_stages=2),
|
||||||
|
triton.Config({"num_warps": 4, "num_stages": 3}, num_warps=4, num_stages=3),
|
||||||
|
triton.Config({"num_warps": 8, "num_stages": 3}, num_warps=8, num_stages=3),
|
||||||
|
],
|
||||||
|
key=[
|
||||||
|
"B",
|
||||||
|
"R",
|
||||||
|
"C",
|
||||||
|
"BLOCK_SIZE",
|
||||||
|
], # Retune if matrix dimensions change significantly
|
||||||
|
)
|
||||||
|
@triton.jit
|
||||||
|
def auction_lap_kernel(
|
||||||
|
cost_ptr,
|
||||||
|
assign_ptr,
|
||||||
|
stride_b,
|
||||||
|
stride_r,
|
||||||
|
stride_c,
|
||||||
|
stride_assign_b,
|
||||||
|
stride_assign_r,
|
||||||
|
B: tl.constexpr,
|
||||||
|
R: tl.constexpr,
|
||||||
|
C: tl.constexpr,
|
||||||
|
epsilon,
|
||||||
|
max_iter,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
|
||||||
|
cost_base = cost_ptr + pid * stride_b
|
||||||
|
assign_base = assign_ptr + pid * stride_assign_b
|
||||||
|
|
||||||
|
offs = tl.arange(0, BLOCK_SIZE)
|
||||||
|
col_mask = offs < C
|
||||||
|
|
||||||
|
# Prices and Owners in SRAM/Registers
|
||||||
|
prices = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
|
||||||
|
owners = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
|
||||||
|
row_to_col = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
|
||||||
|
|
||||||
|
iter_idx = 0
|
||||||
|
unassigned_count = R
|
||||||
|
|
||||||
|
loop_continue = tl.full([], 1, dtype=tl.int1)
|
||||||
|
|
||||||
|
# Loop condition:
|
||||||
|
# 1. unassigned_count > 0: Logic handled inside, but we need a break mechanism
|
||||||
|
# 2. iter_idx < max_iter: Safety break
|
||||||
|
# 3. loop_continue: Did we make progress last time?
|
||||||
|
|
||||||
|
while unassigned_count > 0 and iter_idx < max_iter and loop_continue:
|
||||||
|
# Reset progress flag
|
||||||
|
# loop_continue &= False
|
||||||
|
loop_continue = tl.full([], 0, dtype=tl.int1)
|
||||||
|
|
||||||
|
# Gauss-Seidel pass over all rows
|
||||||
|
for i in tl.range(0, R):
|
||||||
|
# Check if row i is unassigned
|
||||||
|
curr_c = tl.sum(tl.where(offs == i, row_to_col, 0))
|
||||||
|
|
||||||
|
if curr_c == -1:
|
||||||
|
# Load costs
|
||||||
|
row_cost_ptr = cost_base + i * stride_r + offs
|
||||||
|
row_costs = tl.load(row_cost_ptr, mask=col_mask, other=-torch.inf)
|
||||||
|
|
||||||
|
# Net Value
|
||||||
|
values = row_costs - prices
|
||||||
|
|
||||||
|
# Find Best
|
||||||
|
best_val, best_idx = tl.max(values, axis=0, return_indices=True)
|
||||||
|
|
||||||
|
# CRITICAL: Only proceed if this is a valid edge (not -inf)
|
||||||
|
if best_val > -torch.inf:
|
||||||
|
# We have a valid move, so we continue the outer loop
|
||||||
|
loop_continue = tl.full([], 1, dtype=tl.int1)
|
||||||
|
|
||||||
|
# Find Second Best
|
||||||
|
mask_not_best = (offs != best_idx) & col_mask
|
||||||
|
vals_no_best = tl.where(mask_not_best, values, -torch.inf)
|
||||||
|
second_best_val = tl.max(vals_no_best, axis=0)
|
||||||
|
|
||||||
|
# Compute Bid
|
||||||
|
bid = best_val - second_best_val + epsilon
|
||||||
|
|
||||||
|
# Update Price
|
||||||
|
prices = tl.where(offs == best_idx, prices + bid, prices)
|
||||||
|
|
||||||
|
# Update Owners
|
||||||
|
prev_owner = tl.sum(tl.where(offs == best_idx, owners, 0))
|
||||||
|
|
||||||
|
if prev_owner != -1:
|
||||||
|
# Kick out previous owner
|
||||||
|
row_to_col = tl.where(offs == prev_owner, -1, row_to_col)
|
||||||
|
unassigned_count += 1
|
||||||
|
|
||||||
|
# Assign to current row
|
||||||
|
owners = tl.where(offs == best_idx, i, owners)
|
||||||
|
row_to_col = tl.where(offs == i, best_idx, row_to_col)
|
||||||
|
unassigned_count -= 1
|
||||||
|
|
||||||
|
iter_idx += 1
|
||||||
|
|
||||||
|
# Store Result
|
||||||
|
store_offs = tl.arange(0, BLOCK_SIZE)
|
||||||
|
store_mask = store_offs < R
|
||||||
|
tl.store(assign_base + store_offs, row_to_col, mask=store_mask)
|
||||||
|
|
||||||
|
|
||||||
|
# -----------------------------------------------------------------------------
|
||||||
|
# Python Helpers
|
||||||
|
# -----------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def rescale_simple(
|
||||||
|
t: torch.Tensor,
|
||||||
|
target_min: float = 0.0,
|
||||||
|
target_max: float = 1.0,
|
||||||
|
*,
|
||||||
|
start_dim: int = 1,
|
||||||
|
eps: float = 1e-07,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
width = target_max - target_min
|
||||||
|
if width == 0.0:
|
||||||
|
return torch.zeros_like(t)
|
||||||
|
orig_shape = t.shape
|
||||||
|
t = t.flatten(start_dim=start_dim)
|
||||||
|
min_val, max_val = t.aminmax(dim=-1, keepdim=True)
|
||||||
|
normalized = t - min_val
|
||||||
|
normalized /= (max_val - min_val).add_(eps)
|
||||||
|
normalized *= width
|
||||||
|
if target_min != 0.0:
|
||||||
|
normalized += target_min
|
||||||
|
return normalized.clamp_(target_min, target_max).reshape(orig_shape)
|
||||||
|
|
||||||
|
|
||||||
|
def _greedy_fill_missing(assignments: torch.Tensor, C: int) -> None:
|
||||||
|
"""
|
||||||
|
Fills unassigned rows (-1) in the assignments tensor with available columns.
|
||||||
|
This acts as a fallback when the Auction algorithm hits max_iter without
|
||||||
|
full convergence.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
assignments: Tensor of shape (B, R) containing col indices or -1.
|
||||||
|
C: Total number of columns available.
|
||||||
|
"""
|
||||||
|
# Identify which batch items have unassigned rows
|
||||||
|
# This is usually a very small subset (e.g., < 1% of the batch)
|
||||||
|
problem_batches = (assignments == -1).any(dim=1).nonzero().flatten()
|
||||||
|
|
||||||
|
if problem_batches.numel() == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
device = assignments.device
|
||||||
|
|
||||||
|
# Iterate only over the problematic batch items
|
||||||
|
# (Looping is acceptable here as B_subset is typically tiny)
|
||||||
|
for b_idx in problem_batches:
|
||||||
|
# 1. Find which rows are missing an assignment
|
||||||
|
row_mask = assignments[b_idx] == -1
|
||||||
|
missing_rows = row_mask.nonzero().flatten()
|
||||||
|
n_needed = missing_rows.shape[0]
|
||||||
|
|
||||||
|
# 2. Find which columns are already used
|
||||||
|
used_cols = assignments[b_idx][~row_mask]
|
||||||
|
|
||||||
|
# 3. Find free columns (Set difference: All - Used)
|
||||||
|
# Create a boolean mask of all columns, then mark used ones as False
|
||||||
|
# efficient on GPU for mid-sized C
|
||||||
|
col_mask = torch.ones(C, device=device, dtype=torch.bool)
|
||||||
|
col_mask[used_cols.long()] = False
|
||||||
|
|
||||||
|
free_cols = col_mask.nonzero().flatten()
|
||||||
|
|
||||||
|
# 4. Assign the first N free columns to the N missing rows
|
||||||
|
# Since R <= C in this context (due to transpose logic in wrapper),
|
||||||
|
# free_cols.numel() is guaranteed to be >= n_needed.
|
||||||
|
assignments[b_idx, missing_rows] = free_cols[:n_needed].to(assignments.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def batch_linear_assignment(
|
||||||
|
cost_matrix: torch.Tensor,
|
||||||
|
*,
|
||||||
|
maximize: bool = False,
|
||||||
|
max_iter: int | None = None,
|
||||||
|
fill_missing: bool = True,
|
||||||
|
rescale_costs: tuple[float, float] | None = (0.0, 1.0),
|
||||||
|
invert_costs_mode: bool = True,
|
||||||
|
eps: float = 1e-3,
|
||||||
|
):
|
||||||
|
if cost_matrix.ndim != 3:
|
||||||
|
raise ValueError("Cost matrix must be (B, R, C)")
|
||||||
|
if not cost_matrix.is_cuda:
|
||||||
|
raise ValueError("Cost matrix must be a CUDA tensor")
|
||||||
|
if not cost_matrix.is_contiguous():
|
||||||
|
raise ValueError("Cost matrix must be contiguous")
|
||||||
|
|
||||||
|
B, R, C = cost_matrix.shape
|
||||||
|
device = cost_matrix.device
|
||||||
|
|
||||||
|
# 1. Handle Rectangular Matrices
|
||||||
|
# The Auction algorithm assigns Rows -> Cols.
|
||||||
|
# It naturally handles R <= C (finding best col for every row).
|
||||||
|
# If R > C, we must transpose to match Cols -> Rows, then invert the result.
|
||||||
|
if R > C:
|
||||||
|
transposed = True
|
||||||
|
cost_matrix = cost_matrix.mT.contiguous()
|
||||||
|
# Swap R and C for the kernel execution
|
||||||
|
R, C = C, R
|
||||||
|
else:
|
||||||
|
transposed = False
|
||||||
|
|
||||||
|
if rescale_costs is not None:
|
||||||
|
cost_matrix = rescale_simple(cost_matrix, *rescale_costs)
|
||||||
|
# Note: We use float32 for atomic compatibility and speed
|
||||||
|
cost_matrix = cost_matrix.to(torch.float32, copy=rescale_costs is None)
|
||||||
|
|
||||||
|
if not maximize:
|
||||||
|
# Maximize (Value - Price) -> Minimize Cost
|
||||||
|
cost_matrix = cost_matrix.neg_()
|
||||||
|
if invert_costs_mode and rescale_costs is not None:
|
||||||
|
cost_matrix += sum(rescale_costs)
|
||||||
|
|
||||||
|
assignments = torch.full(
|
||||||
|
(B, R),
|
||||||
|
-1,
|
||||||
|
device=device,
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
|
||||||
|
max_dim = max(R, C)
|
||||||
|
BLOCK_SIZE = max(32, triton.next_power_of_2(max_dim))
|
||||||
|
|
||||||
|
# Safety limit
|
||||||
|
max_iter = max_iter if max_iter is not None else int(max(2000, R * C))
|
||||||
|
|
||||||
|
grid = (B,)
|
||||||
|
|
||||||
|
auction_lap_kernel[grid](
|
||||||
|
cost_matrix,
|
||||||
|
assignments,
|
||||||
|
cost_matrix.stride(0),
|
||||||
|
cost_matrix.stride(1),
|
||||||
|
cost_matrix.stride(2),
|
||||||
|
assignments.stride(0),
|
||||||
|
assignments.stride(1),
|
||||||
|
B,
|
||||||
|
R,
|
||||||
|
C,
|
||||||
|
eps,
|
||||||
|
max_iter,
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
if fill_missing:
|
||||||
|
_greedy_fill_missing(assignments, C)
|
||||||
|
|
||||||
|
assignments = assignments.long()
|
||||||
|
|
||||||
|
if not transposed:
|
||||||
|
return assignments
|
||||||
|
|
||||||
|
# 2. Post-process Rectangular Results
|
||||||
|
# We computed Col -> Row. We need Row -> Col.
|
||||||
|
# assignments shape is currently (B, Original_Cols)
|
||||||
|
# We want output shape (B, Original_Rows)
|
||||||
|
|
||||||
|
real_rows = C # C is the 'large' dimension (Original Rows)
|
||||||
|
output = torch.full((B, real_rows), -1, device=device, dtype=torch.long)
|
||||||
|
|
||||||
|
# Create indices for the scatter source
|
||||||
|
# We want: output[row_idx] = col_idx
|
||||||
|
# Currently we have: assignments[col_idx] = row_idx
|
||||||
|
src_col_indices = torch.arange(R, device=device).unsqueeze(0).expand(B, R)
|
||||||
|
|
||||||
|
# We use scatter. index=assignments (the rows), src=col_indices
|
||||||
|
# To handle -1s in assignments, we clamp to 0 and then mask the result
|
||||||
|
safe_assigns = assignments.clamp(min=0)
|
||||||
|
output.scatter_(1, safe_assigns, src_col_indices)
|
||||||
|
|
||||||
|
# Cleanup: Any row that wasn't targeted by the scatter should be -1
|
||||||
|
# The scatter might have written to index 0 if assignment was -1
|
||||||
|
# Re-verify logic:
|
||||||
|
for b in range(B):
|
||||||
|
valid_mask = assignments[b] >= 0
|
||||||
|
# Reset output
|
||||||
|
output[b].fill_(-1)
|
||||||
|
# Only write valid mappings
|
||||||
|
# output[b, row_id] = col_id
|
||||||
|
output[b, assignments[b, valid_mask]] = src_col_indices[b, valid_mask]
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def assignments_to_indices(
|
||||||
|
assignments: torch.Tensor,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Converts a dense assignment tensor (from Triton/Hungraian) to
|
||||||
|
batched SciPy-style indices.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
assignments (torch.Tensor): Shape (B, R). Values are col indices or -1.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
row_ind (torch.Tensor): Shape (B, K) where K = min(R, C).
|
||||||
|
col_ind (torch.Tensor): Shape (B, K).
|
||||||
|
"""
|
||||||
|
B, R = assignments.shape
|
||||||
|
device = assignments.device
|
||||||
|
|
||||||
|
# 1. Create a mask of valid assignments (values >= 0)
|
||||||
|
# In a rectangular assignment, the number of valid matches
|
||||||
|
# is always min(Rows, Cols).
|
||||||
|
mask = assignments >= 0
|
||||||
|
|
||||||
|
# 2. Extract Column Indices
|
||||||
|
# We select the values from the assignment tensor that are valid.
|
||||||
|
# We reshape to (B, -1) to preserve the batch dimension.
|
||||||
|
col_ind = assignments[mask].view(B, -1)
|
||||||
|
|
||||||
|
# 3. Extract Row Indices
|
||||||
|
# We need a grid of row indices [0, 1, 2, ... R-1] repeated B times
|
||||||
|
row_grid = (
|
||||||
|
torch.arange(R, device=device, dtype=assignments.dtype)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(B, R)
|
||||||
|
)
|
||||||
|
row_ind = row_grid[mask].view(B, -1)
|
||||||
|
|
||||||
|
return row_ind, col_ind
|
||||||
|
|
||||||
|
|
||||||
|
def batch_linear_assignment_shuffled(
|
||||||
|
cost_matrix: torch.Tensor,
|
||||||
|
*args,
|
||||||
|
**kwargs: dict,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
generator = kwargs.pop("generator", None)
|
||||||
|
# cost_matrix: [B, R, C]
|
||||||
|
B, R = cost_matrix.shape[:2]
|
||||||
|
|
||||||
|
# 1. Generate a random permutation for the rows
|
||||||
|
# We use one perm for the whole batch for efficiency,
|
||||||
|
# or you can do it per-batch-item if B is small and quality is critical.
|
||||||
|
# Here we shuffle all rows commonly.
|
||||||
|
perm = torch.randperm(
|
||||||
|
R,
|
||||||
|
device=cost_matrix.device,
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Shuffle the input (Row dimension is dim 1)
|
||||||
|
# This creates a shuffled view/copy of the cost matrix
|
||||||
|
shuffled_cost = cost_matrix[:, perm, :]
|
||||||
|
|
||||||
|
# 3. Run the Solver
|
||||||
|
shuffled_assignments = batch_linear_assignment(
|
||||||
|
shuffled_cost,
|
||||||
|
*args,
|
||||||
|
**kwargs,
|
||||||
|
) # Returns [B, R]
|
||||||
|
|
||||||
|
# 4. Un-shuffle the results
|
||||||
|
# We need to map the results back to their original row positions.
|
||||||
|
# shuffled_assignments[b, i] corresponds to the row 'perm[i]'
|
||||||
|
# We want final_assignments[b, perm[i]] = shuffled_assignments[b, i]
|
||||||
|
|
||||||
|
# Create the inverse permutation or just scatter back
|
||||||
|
final_assignments = torch.empty_like(shuffled_assignments)
|
||||||
|
|
||||||
|
# Expand perm for the batch: [B, R]
|
||||||
|
batch_perm = perm.unsqueeze(0).expand(B, R)
|
||||||
|
|
||||||
|
# Scatter the results back to original positions
|
||||||
|
# dim=1, index=batch_perm, src=shuffled_assignments
|
||||||
|
final_assignments.scatter_(1, batch_perm, shuffled_assignments)
|
||||||
|
|
||||||
|
return final_assignments
|
||||||
+101
-168
@@ -1,186 +1,119 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from comfy.k_diffusion.sampling import to_d
|
from comfy.k_diffusion.sampling import to_d
|
||||||
|
|
||||||
# def scale_noise_(
|
F = torch.nn.functional
|
||||||
# noise,
|
|
||||||
# factor=1.0,
|
|
||||||
# *,
|
|
||||||
# normalized=True,
|
|
||||||
# normalize_dims=(-3, -2, -1),
|
|
||||||
# ):
|
|
||||||
# if not normalized or noise.numel() == 0:
|
|
||||||
# return noise.mul_(factor) if factor != 1 else noise
|
|
||||||
# mean, std = (
|
|
||||||
# noise.mean(dim=normalize_dims, keepdim=True),
|
|
||||||
# noise.std(dim=normalize_dims, keepdim=True),
|
|
||||||
# )
|
|
||||||
# return latent.normalize_to_scale(
|
|
||||||
# noise.sub_(mean).div_(std).clamp(-1, 1), -1.0, 1.0, dim=normalize_dims
|
|
||||||
# ).mul_(factor)
|
|
||||||
|
|
||||||
|
|
||||||
# def scale_noise(
|
|
||||||
# noise,
|
|
||||||
# factor=1.0,
|
|
||||||
# *,
|
|
||||||
# normalized=True,
|
|
||||||
# normalize_dims=(-3, -2, -1),
|
|
||||||
# ):
|
|
||||||
# if not normalized or noise.numel() == 0:
|
|
||||||
# return noise * factor if factor != 1 else noise
|
|
||||||
# mean, std = (
|
|
||||||
# noise.mean(dim=normalize_dims, keepdim=True),
|
|
||||||
# noise.std(dim=normalize_dims, keepdim=True),
|
|
||||||
# )
|
|
||||||
# return (noise - mean).div_(std).mul_(factor)
|
|
||||||
|
|
||||||
|
|
||||||
def scale_noise(
|
def scale_noise(
|
||||||
noise,
|
noise: torch.Tensor,
|
||||||
factor=1.0,
|
factor: float = 1.0,
|
||||||
*,
|
*,
|
||||||
normalized=True,
|
normalized: bool = True,
|
||||||
normalize_dims=(-3, -2, -1),
|
normalize_dims: tuple[int, ...] | None = None,
|
||||||
):
|
eps: float | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if factor == 0:
|
||||||
|
return torch.zeros_like(noise)
|
||||||
if not normalized or noise.numel() == 0:
|
if not normalized or noise.numel() == 0:
|
||||||
return noise * factor if factor != 1 else noise
|
return noise * factor if factor != 1 else noise
|
||||||
noise = noise / noise.std(dim=normalize_dims, keepdim=True)
|
if eps is None:
|
||||||
return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor)
|
eps = torch.finfo(noise.dtype).eps * 1.25
|
||||||
|
if normalize_dims is None:
|
||||||
|
normalize_dims = tuple(
|
||||||
def _quantile_norm_scaledown(
|
range(
|
||||||
noise: torch.Tensor,
|
max(0, min(1, noise.ndim - 1)),
|
||||||
nq: torch.Tensor,
|
noise.ndim,
|
||||||
**_kwargs: dict,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
mv = noise.abs().max().detach().item()
|
|
||||||
return noise if mv == 0 else torch.where(noise.abs() > nq, noise * (nq / mv), noise)
|
|
||||||
|
|
||||||
|
|
||||||
quantile_handlers = {
|
|
||||||
"clamp": lambda noise, nq, **_kwargs: noise.clamp(-nq, nq),
|
|
||||||
"scale_down": _quantile_norm_scaledown,
|
|
||||||
"tanh": lambda noise, nq, **_kwargs: noise.tanh().mul_(nq.abs()),
|
|
||||||
"tanh_outliers": lambda noise, nq, **_kwargs: torch.where(
|
|
||||||
noise.abs() > nq,
|
|
||||||
noise.tanh().mul_(nq.abs()),
|
|
||||||
noise,
|
|
||||||
),
|
|
||||||
"sigmoid": lambda noise, nq, **_kwargs: noise.sigmoid()
|
|
||||||
.mul_(nq.abs())
|
|
||||||
.copysign(noise),
|
|
||||||
"sigmoid_outliers": lambda noise, nq, **_kwargs: torch.where(
|
|
||||||
noise.abs() > nq,
|
|
||||||
noise.sigmoid().mul_(nq.abs()).copysign(noise),
|
|
||||||
noise,
|
|
||||||
),
|
|
||||||
"tenth": lambda noise, nq, **_kwargs: torch.where(
|
|
||||||
noise.abs() > nq,
|
|
||||||
noise * 0.1,
|
|
||||||
noise,
|
|
||||||
),
|
|
||||||
"half": lambda noise, nq, **_kwargs: torch.where(
|
|
||||||
noise.abs() > nq,
|
|
||||||
noise * 0.5,
|
|
||||||
noise,
|
|
||||||
),
|
|
||||||
"zero": lambda noise, nq, **_kwargs: torch.where(noise.abs() > nq, 0, noise),
|
|
||||||
"reverse_zero": lambda noise, nq, **_kwargs: torch.where(
|
|
||||||
noise.abs() >= nq,
|
|
||||||
noise,
|
|
||||||
0,
|
|
||||||
),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
# Initial version based on Studentt distribution normalizatino from https://github.com/Clybius/ComfyUI-Extra-Samplers/
|
|
||||||
def quantile_normalize(
|
|
||||||
noise: torch.Tensor,
|
|
||||||
*,
|
|
||||||
quantile: float = 0.75,
|
|
||||||
dim: int | None = 1,
|
|
||||||
flatten: bool = True,
|
|
||||||
nq_fac: float = 1.0,
|
|
||||||
pow_fac: float = 0.5,
|
|
||||||
strategy: str = "clamp",
|
|
||||||
strategy_handler=None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if quantile is None or quantile <= 0 or quantile >= 1:
|
|
||||||
return noise
|
|
||||||
orig_shape = noise.shape
|
|
||||||
if isinstance(quantile, (tuple, list)):
|
|
||||||
quantile = torch.tensor(
|
|
||||||
quantile,
|
|
||||||
device=noise.device,
|
|
||||||
dtype=noise.dtype,
|
|
||||||
)
|
|
||||||
qdim = dim
|
|
||||||
if noise.ndim > 1 and flatten:
|
|
||||||
if qdim is not None and qdim >= noise.ndim:
|
|
||||||
qdim = 1 if noise.ndim > 2 else None
|
|
||||||
if qdim is None:
|
|
||||||
flatdim = 0
|
|
||||||
elif -1 < qdim < 2: # 0, 1
|
|
||||||
flatdim = qdim + 1
|
|
||||||
elif 1 < qdim < 4: # 2, 3
|
|
||||||
noise = noise.movedim(qdim, 1)
|
|
||||||
tempshape = noise.shape
|
|
||||||
flatdim = 2
|
|
||||||
else:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot handling quantile normalization flattening dims > 3",
|
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
flatdim = None
|
|
||||||
nq = torch.quantile(
|
|
||||||
(noise if flatdim is None else noise.flatten(start_dim=flatdim)).abs(),
|
|
||||||
quantile,
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim)
|
|
||||||
nq = nq.mul_(nq_fac).reshape(*nq_shape)
|
|
||||||
handler = (
|
|
||||||
quantile_handlers.get(strategy)
|
|
||||||
if strategy_handler is None
|
|
||||||
else strategy_handler
|
|
||||||
)
|
|
||||||
if handler is None:
|
|
||||||
raise ValueError("Unknown strategy")
|
|
||||||
noise = handler(
|
|
||||||
noise,
|
|
||||||
nq,
|
|
||||||
dim=dim,
|
|
||||||
flatten=flatten,
|
|
||||||
)
|
|
||||||
noise = noise.abs().pow(pow_fac).copysign(noise)
|
|
||||||
if flatdim is not None and qdim in {2, 3}:
|
|
||||||
return (
|
|
||||||
noise.reshape(tempshape).movedim(1, qdim).reshape(orig_shape).contiguous()
|
|
||||||
)
|
)
|
||||||
return noise
|
std, mean = torch.std_mean(noise, dim=normalize_dims, keepdim=True)
|
||||||
|
noise = noise - mean
|
||||||
|
if factor != 1:
|
||||||
|
std /= factor
|
||||||
|
return noise.div_(std.clamp_min_(eps) if factor >= 0 else std.clamp_max_(-eps))
|
||||||
|
|
||||||
|
|
||||||
# def scale_noise(
|
def range_wrap(
|
||||||
# noise,
|
x: torch.Tensor,
|
||||||
# factor=1.0,
|
min_val: float | torch.Tensor,
|
||||||
# *,
|
max_val: float | torch.Tensor,
|
||||||
# normalized=True,
|
) -> torch.Tensor:
|
||||||
# normalize_dims=(-3, -2, -1),
|
return min_val + (x - min_val).remainder_(max_val - min_val)
|
||||||
# ):
|
|
||||||
# if not normalized or noise.numel() == 0:
|
|
||||||
# return noise.mul_(factor) if factor != 1 else noise
|
def softplus_soft_clamp(
|
||||||
# n = (
|
t: torch.Tensor,
|
||||||
# torch.nn.LayerNorm(noise.shape[1:])
|
min_val: torch.Tensor | float = 0.0,
|
||||||
# if normalize_dims == (-3, -2, -1)
|
max_val: torch.Tensor | float = 1.0,
|
||||||
# else torch.nn.InstanceNorm2d(noise.shape[1])
|
*,
|
||||||
# ).to(noise)
|
# We define stiffness as a multiplier (beta) for the softplus function.
|
||||||
# return n(noise) * factor
|
# Higher stiffness = sharper transition.
|
||||||
# return latent.normalize_to_scale(
|
stiffness: float = 1.0,
|
||||||
# n(noise).clamp_(-1, 1), -1, 1, dim=normalize_dims
|
safe: bool = True,
|
||||||
# ).mul_(factor)
|
) -> torch.Tensor:
|
||||||
|
if isinstance(min_val, (float, int)):
|
||||||
|
min_val = t.new_tensor(min_val)
|
||||||
|
if isinstance(max_val, (float, int)):
|
||||||
|
max_val = t.new_tensor(max_val)
|
||||||
|
|
||||||
|
if stiffness < 1e-04:
|
||||||
|
return t.clamp(min_val, max_val)
|
||||||
|
|
||||||
|
# Calculate how much we are exceeding the Max
|
||||||
|
# softplus(beta * x) / beta
|
||||||
|
upper_overshoot = F.softplus((t - max_val).mul_(stiffness)).div_(-stiffness)
|
||||||
|
|
||||||
|
# Calculate how much we are falling short of the Min
|
||||||
|
lower_undershoot = F.softplus((min_val - t).mul_(stiffness)).div_(stiffness)
|
||||||
|
|
||||||
|
# Apply corrections:
|
||||||
|
# Original - (Amount over max) + (Amount under min)
|
||||||
|
t = upper_overshoot.add_(t).add_(lower_undershoot)
|
||||||
|
if safe:
|
||||||
|
t = t.clamp(min_val, max_val)
|
||||||
|
return t
|
||||||
|
|
||||||
|
|
||||||
|
def flip_tensor_range(
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
min_neg: torch.Tensor | None = None,
|
||||||
|
max_pos: torch.Tensor | None = None,
|
||||||
|
return_ranges: bool = False,
|
||||||
|
dim: int = -1,
|
||||||
|
eps: float | None = None,
|
||||||
|
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
if eps is None:
|
||||||
|
eps = torch.finfo(x.dtype).eps * 1.25
|
||||||
|
# 1. Use the provided maximum positive values, or calculate them dynamically
|
||||||
|
if max_pos is None:
|
||||||
|
max_pos = (
|
||||||
|
torch.clamp_min(x, 0.0).max(dim=dim, keepdim=True).values.clamp_min_(eps)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Use the provided minimum negative values, or calculate them dynamically
|
||||||
|
if min_neg is None:
|
||||||
|
min_neg = (
|
||||||
|
torch.clamp_max(x, 0.0).min(dim=dim, keepdim=True).values.clamp_max_(-eps)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 3. Separate positive and negative elements
|
||||||
|
is_pos = x >= 0
|
||||||
|
|
||||||
|
# 4. Flip positive side: [0, max_pos] -> [eps, max_pos + eps]
|
||||||
|
x_pos = x.clamp_min(eps)
|
||||||
|
flipped_pos = (max_pos + eps) - x_pos
|
||||||
|
|
||||||
|
# 5. Flip negative side: [min_neg, 0] -> [min_neg - eps, -eps]
|
||||||
|
x_neg = x.clamp_max(-eps)
|
||||||
|
flipped_neg = (min_neg - eps) - x_neg
|
||||||
|
|
||||||
|
# 6. Recombine the domains
|
||||||
|
result = torch.where(is_pos, flipped_pos, flipped_neg)
|
||||||
|
return (result, max_pos, min_neg) if return_ranges else result
|
||||||
|
|
||||||
|
|
||||||
def find_first_unsorted(tensor, desc=True):
|
def find_first_unsorted(tensor, desc=True):
|
||||||
|
|||||||
@@ -0,0 +1,238 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Callable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from .utils import fallback
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
try:
|
||||||
|
import pytorch_wavelets as ptwav
|
||||||
|
import pywt
|
||||||
|
|
||||||
|
HAVE_WAVELETS = True
|
||||||
|
except ImportError:
|
||||||
|
ptwav = None
|
||||||
|
pywt = None
|
||||||
|
HAVE_WAVELETS = False
|
||||||
|
|
||||||
|
|
||||||
|
class Wavelet:
|
||||||
|
DEFAULT_MODE = "symmetric"
|
||||||
|
DEFAULT_LEVEL = 3
|
||||||
|
DEFAULT_WAVE = "db4"
|
||||||
|
DEFAULT_USE_1D_DWT = False
|
||||||
|
DEFAULT_USE_DTCWT = False
|
||||||
|
DEFAULT_QSHIFT = "qshift_a"
|
||||||
|
DEFAULT_BIORT = "near_sym_a"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
wave: str = DEFAULT_WAVE,
|
||||||
|
level: int = DEFAULT_LEVEL,
|
||||||
|
mode: str = DEFAULT_MODE,
|
||||||
|
use_1d_dwt: bool = DEFAULT_USE_1D_DWT,
|
||||||
|
use_dtcwt: bool = DEFAULT_USE_DTCWT,
|
||||||
|
biort: str = DEFAULT_BIORT,
|
||||||
|
qshift: str = DEFAULT_QSHIFT,
|
||||||
|
inv_wave: str | None = None,
|
||||||
|
inv_mode: str | None = None,
|
||||||
|
inv_biort: str | None = None,
|
||||||
|
inv_qshift=None,
|
||||||
|
device: str | torch.device | None = None,
|
||||||
|
):
|
||||||
|
if not HAVE_WAVELETS:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Wavelet noise requires the pytorch_wavelets package to be installed in your Python environment",
|
||||||
|
)
|
||||||
|
inv_wave = fallback(inv_wave, wave)
|
||||||
|
inv_mode = fallback(inv_mode, mode)
|
||||||
|
inv_biort = fallback(inv_biort, biort)
|
||||||
|
inv_qshift = fallback(inv_qshift, qshift)
|
||||||
|
if use_dtcwt:
|
||||||
|
fwdfun, invfun = ptwav.DTCWTForward, ptwav.DTCWTInverse
|
||||||
|
elif use_1d_dwt:
|
||||||
|
fwdfun, invfun = ptwav.DWT1DForward, ptwav.DWT1DInverse
|
||||||
|
else:
|
||||||
|
fwdfun, invfun = ptwav.DWTForward, ptwav.DWTInverse
|
||||||
|
if use_dtcwt:
|
||||||
|
self._wavelet_forward = fwdfun(
|
||||||
|
J=level,
|
||||||
|
mode=mode,
|
||||||
|
biort=biort,
|
||||||
|
qshift=qshift,
|
||||||
|
)
|
||||||
|
self._wavelet_inverse = invfun(
|
||||||
|
mode=inv_mode,
|
||||||
|
biort=inv_biort,
|
||||||
|
qshift=inv_qshift,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._wavelet_forward = fwdfun(J=level, wave=wave, mode=mode)
|
||||||
|
self._wavelet_inverse = invfun(wave=inv_wave, mode=inv_mode)
|
||||||
|
if device is not None:
|
||||||
|
self._wavelet_forward = self._wavelet_forward.to(device=device)
|
||||||
|
self._wavelet_inverse = self._wavelet_inverse.to(device=device)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
t: torch.Tensor,
|
||||||
|
*,
|
||||||
|
forward_function: Callable | None = None,
|
||||||
|
) -> tuple[torch.Tensor, tuple]:
|
||||||
|
return fallback(forward_function, self._wavelet_forward)(t)
|
||||||
|
|
||||||
|
def inverse(
|
||||||
|
self,
|
||||||
|
yl: torch.Tensor,
|
||||||
|
yh: tuple,
|
||||||
|
*,
|
||||||
|
inverse_function: Callable | None = None,
|
||||||
|
two_step_inverse: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
inverse_function = fallback(inverse_function, self._wavelet_inverse)
|
||||||
|
if not two_step_inverse:
|
||||||
|
return inverse_function((yl, yh))
|
||||||
|
result = inverse_function((torch.zeros_like(yl), yh))
|
||||||
|
result += inverse_function((
|
||||||
|
yl,
|
||||||
|
tuple(torch.zeros_like(yh_band) for yh_band in yh),
|
||||||
|
))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def to(self, *args: list, copy: bool = False, **kwargs: dict) -> Wavelet:
|
||||||
|
o = Wavelet.__new__(Wavelet) if copy else self
|
||||||
|
o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001
|
||||||
|
o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001
|
||||||
|
return o
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def wavelist() -> tuple:
|
||||||
|
return tuple(pywt.wavelist()) if HAVE_WAVELETS else ()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def biortlist() -> tuple:
|
||||||
|
return (
|
||||||
|
("near_sym_a", "near_sym_b", "antonini", "legall") if HAVE_WAVELETS else ()
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def qshiftlist() -> tuple:
|
||||||
|
return (
|
||||||
|
("qshift_a", "qshift_b", "qshift_c", "qshift_d", "qshift_06")
|
||||||
|
if HAVE_WAVELETS
|
||||||
|
else ()
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def modelist() -> tuple:
|
||||||
|
return (
|
||||||
|
(
|
||||||
|
"symmetric",
|
||||||
|
"zero",
|
||||||
|
"reflect",
|
||||||
|
"replicate",
|
||||||
|
"periodization",
|
||||||
|
"periodic",
|
||||||
|
"constant",
|
||||||
|
)
|
||||||
|
if HAVE_WAVELETS
|
||||||
|
else ()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def expand_yh_scales(
|
||||||
|
yh: Sequence,
|
||||||
|
*,
|
||||||
|
yh_scales: float | Sequence = 1.0,
|
||||||
|
) -> float | tuple:
|
||||||
|
yhlen = len(yh)
|
||||||
|
yh_shape = yh[0].shape
|
||||||
|
# Doesn't make sense to target orientations for 1D DWD (3D here).
|
||||||
|
olen = yh_shape[2] if len(yh_shape) > 3 else 1
|
||||||
|
# print(f"\nSIZES: yhlen={yhlen}, olen={olen}, yh_shape={yh[0].shape}")
|
||||||
|
if isinstance(yh_scales, (float, int)):
|
||||||
|
return ((float(yh_scales),) * olen,) * yhlen
|
||||||
|
otemplate = (1.0,) * olen
|
||||||
|
yh_scales = tuple(
|
||||||
|
(float(band),) * olen
|
||||||
|
if isinstance(band, (float, int))
|
||||||
|
else (
|
||||||
|
(
|
||||||
|
*(float(i) for i in band[:olen]),
|
||||||
|
*otemplate[: olen - len(band[:olen])],
|
||||||
|
)
|
||||||
|
if isinstance(band, (tuple, list))
|
||||||
|
else band
|
||||||
|
)
|
||||||
|
for band in yh_scales
|
||||||
|
)
|
||||||
|
if "fill" in yh_scales:
|
||||||
|
fillidx = yh_scales.index("fill")
|
||||||
|
if "fill" in yh_scales[fillidx + 1 :]:
|
||||||
|
raise ValueError("Only one fill allowed.")
|
||||||
|
if fillidx == 0 or len(yh_scales) < 2:
|
||||||
|
raise ValueError(
|
||||||
|
"Invalid fill value, cannot be in the first position or the only item.",
|
||||||
|
)
|
||||||
|
yhslen = len(yh_scales)
|
||||||
|
if yhslen - 1 < yhlen:
|
||||||
|
# Need to pad.
|
||||||
|
fill = (yh_scales[fillidx - 1],) * (yhlen - (len(yh_scales) - 1))
|
||||||
|
yh_scales = (*yh_scales[:fillidx], *fill, *yh_scales[fillidx + 1 :])
|
||||||
|
else:
|
||||||
|
# Just remove the "fill".
|
||||||
|
yh_scales = (*yh_scales[:fillidx], *yh_scales[fillidx + 1 :])
|
||||||
|
return yh_scales[:yhlen]
|
||||||
|
|
||||||
|
|
||||||
|
def wavelet_scaling(
|
||||||
|
yl: torch.Tensor,
|
||||||
|
yh: Sequence,
|
||||||
|
yl_scale: float | torch.Tensor,
|
||||||
|
yh_scales: float | Sequence | None,
|
||||||
|
*,
|
||||||
|
in_place: bool = False,
|
||||||
|
) -> tuple:
|
||||||
|
if not in_place:
|
||||||
|
yl = yl.clone()
|
||||||
|
yh = tuple(yhband.clone() for yhband in yh)
|
||||||
|
if yl_scale != 1.0:
|
||||||
|
yl *= yl_scale
|
||||||
|
yh_scales = expand_yh_scales(
|
||||||
|
yh,
|
||||||
|
yh_scales=yh_scales if yh_scales is not None else 1.0,
|
||||||
|
)
|
||||||
|
for hscale, ht in zip(yh_scales, yh):
|
||||||
|
if isinstance(hscale, (int, float)):
|
||||||
|
ht *= hscale # noqa: PLW2901
|
||||||
|
continue
|
||||||
|
for lidx in range(min(ht.shape[2], len(hscale))):
|
||||||
|
ht[:, :, lidx] *= hscale[lidx]
|
||||||
|
return (yl, yh)
|
||||||
|
|
||||||
|
|
||||||
|
def wavelet_blend(
|
||||||
|
a: tuple,
|
||||||
|
b: tuple,
|
||||||
|
*,
|
||||||
|
yl_factor: torch.Tensor | float,
|
||||||
|
blend_function: Callable,
|
||||||
|
yh_factor: torch.Tensor | float | None = None,
|
||||||
|
yh_blend_function: Callable | None = None,
|
||||||
|
) -> tuple:
|
||||||
|
if not isinstance(yl_factor, torch.Tensor):
|
||||||
|
yl_factor = a[0].new_full((1,), yl_factor)
|
||||||
|
if yh_factor is None:
|
||||||
|
yh_factor = yl_factor
|
||||||
|
elif not isinstance(yh_factor, torch.Tensor):
|
||||||
|
yh_factor = a[0].new_full((1,), yh_factor)
|
||||||
|
yh_blend_function = fallback(yh_blend_function, blend_function)
|
||||||
|
return (
|
||||||
|
blend_function(a[0], b[0], yl_factor),
|
||||||
|
tuple(yh_blend_function(ta, tb, yh_factor) for ta, tb in zip(a[1], b[1])),
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user