Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
79c7b24d1e | ||
|
|
5ea059bed5 | ||
|
|
ee59df94e3 |
@@ -11,5 +11,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
|
||||
"OCS ApplyFilterLatent": nodes.ApplyFilterLatent,
|
||||
"OCS ApplyFilterImage": nodes.ApplyFilterImage,
|
||||
"OCS ExpressionFilteredLatentOperation": nodes.ExpressionFilteredLatentOperationNode,
|
||||
"OCS ExpressionFilteredModelPatch": nodes.ExpressionFilteredModelPatchNode,
|
||||
} | custom_noise.NODE_CLASS_MAPPINGS
|
||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import abc
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from typing import Callable, Any
|
||||
|
||||
from ..external import IntegratedNode
|
||||
from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE
|
||||
from ..noise import scale_noise
|
||||
from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT
|
||||
|
||||
|
||||
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 .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__ = (
|
||||
"Arg",
|
||||
|
||||
+27
-12
@@ -1,20 +1,22 @@
|
||||
import re
|
||||
import operator
|
||||
import re
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from .parser import Parser, ParserSpec, ParseError
|
||||
from .parser import ParseError, Parser, ParserSpec
|
||||
from .types import (
|
||||
Empty,
|
||||
ExpBase,
|
||||
ExpOp,
|
||||
ExpBinOp,
|
||||
ExpSym,
|
||||
ExpStatements,
|
||||
ExpFunAp,
|
||||
ExpTuple,
|
||||
ExpDict,
|
||||
ExpFunAp,
|
||||
ExpKV,
|
||||
ExpMethodAp,
|
||||
ExpOp,
|
||||
ExpReturn,
|
||||
ExpStatements,
|
||||
ExpSym,
|
||||
ExpTuple,
|
||||
)
|
||||
|
||||
COMMA_PRECEDENCE = 2
|
||||
@@ -36,10 +38,11 @@ class Expression:
|
||||
| :> # Key value binop
|
||||
| := # Assignment
|
||||
| ; # Sequencing
|
||||
| :: # Method call
|
||||
| [?:] # Ternary
|
||||
| \[ | ] # Index
|
||||
| \.\.\. # Index ellipsis
|
||||
| '[-\w.]+ # Symbol
|
||||
| '[-\w.:=]+ # Symbol
|
||||
| `?[a-z][\w.]*`? # Function/variable names
|
||||
)
|
||||
\s*
|
||||
@@ -63,7 +66,10 @@ class Expression:
|
||||
tqdm.write(f"* OCS: EVAL: {self.expr}")
|
||||
if not isinstance(self.expr, ExpBase):
|
||||
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):
|
||||
return len(self.expr)
|
||||
@@ -157,9 +163,9 @@ class ExprParserSpec(ParserSpec):
|
||||
def split_funap_args(toks):
|
||||
if not isinstance(toks, (list, tuple)):
|
||||
return ExpTuple((toks,)), 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)
|
||||
})
|
||||
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)}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def null_constant(p, token, bp):
|
||||
@@ -200,6 +206,14 @@ class ExprParserSpec(ParserSpec):
|
||||
p.expect(")")
|
||||
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
|
||||
def left_comma(p, token, left, bp):
|
||||
if p.token == ")":
|
||||
@@ -251,6 +265,7 @@ class ExprParserSpec(ParserSpec):
|
||||
def populate(self):
|
||||
self.add_left(31, self.left_funcall, ("(",))
|
||||
self.add_left(31, self.left_index, ("[",))
|
||||
self.add_left(31, self.left_methodcall, ("::",))
|
||||
self.add_leftright(29, self.left_binop, ("**",))
|
||||
self.add_null(27, self.null_prefixop, ("+", "-", "!"))
|
||||
self.add_left(25, self.left_binop, ("*", "/"))
|
||||
|
||||
+105
-15
@@ -1,8 +1,11 @@
|
||||
import operator
|
||||
import traceback
|
||||
|
||||
from .validation import ValidateArg, Arg, ValidateError
|
||||
from .types import Empty, ExpDict, ExpOp
|
||||
from tqdm import tqdm
|
||||
|
||||
from .types import Empty, ExpDict, ExpOp, ExpReturn, ExpTuple
|
||||
from .util import torch
|
||||
from .validation import Arg, ValidateArg, ValidateError
|
||||
|
||||
|
||||
class HandlerError(Exception):
|
||||
@@ -61,9 +64,12 @@ class BaseHandler:
|
||||
def __call__(self, obj, *, getter):
|
||||
try:
|
||||
val = self.handle(obj, getter)
|
||||
return self.validate_output(obj, val)
|
||||
except ExpReturn:
|
||||
raise
|
||||
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):
|
||||
str_key = isinstance(key, str)
|
||||
@@ -256,6 +262,22 @@ class UnarySimpleMathHandler(SimpleMathHandler):
|
||||
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):
|
||||
input_validators = (Arg.string("name"),)
|
||||
|
||||
@@ -328,6 +350,13 @@ class MaxHandler(MinHandler):
|
||||
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):
|
||||
input_validators = (Arg.present("__callable"),)
|
||||
|
||||
@@ -338,7 +367,7 @@ class UnsafeCallHandler(BaseHandler):
|
||||
)
|
||||
fun = self.safe_get("__callable", obj, getter)
|
||||
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)))
|
||||
kwargs = {k: self.safe_get(k, obj, getter) for k in obj.kwargs}
|
||||
return fun(*args, **kwargs)
|
||||
@@ -365,6 +394,53 @@ class SetVarHandler(BaseHandler):
|
||||
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 = {
|
||||
"||": OrHandler(),
|
||||
"&&": AndHandler(),
|
||||
@@ -385,21 +461,26 @@ for k, alias in (
|
||||
|
||||
|
||||
MATH_HANDLERS = {
|
||||
"*": SimpleMathHandler(operator.mul),
|
||||
"**": SimpleMathHandler(operator.pow),
|
||||
"+": SimpleMathHandler(operator.add),
|
||||
"-": MinusHandler(),
|
||||
"*": SimpleMathHandler(operator.mul),
|
||||
"/": SimpleMathHandler(operator.truediv),
|
||||
"//": SimpleMathHandler(operator.floordiv),
|
||||
"**": SimpleMathHandler(operator.pow),
|
||||
"mod": SimpleMathHandler(operator.mod),
|
||||
"neg": UnarySimpleMathHandler(operator.neg),
|
||||
"between": BetweenHandler(),
|
||||
"<": RelComparisonHandler(operator.lt),
|
||||
"<=": RelComparisonHandler(operator.le),
|
||||
">": RelComparisonHandler(operator.gt),
|
||||
">=": 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(),
|
||||
"min": MinHandler(),
|
||||
"mod": SimpleMathHandler(operator.mod),
|
||||
"neg": UnarySimpleMathHandler(operator.neg),
|
||||
"sum": SumHandler(),
|
||||
}
|
||||
for k, alias in (
|
||||
("+", "add"),
|
||||
@@ -412,14 +493,23 @@ for k, alias in (
|
||||
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
|
||||
|
||||
MISC_HANDLERS = {
|
||||
"is_set": IsSetHandler(),
|
||||
"and": SimpleOpHandler(operator.and_),
|
||||
"comment": CommentHandler(),
|
||||
"concat": SimpleOpHandler(operator.concat),
|
||||
"contains": SimpleOpHandler(operator.contains),
|
||||
"dict": DictHandler(),
|
||||
"get": GetHandler(),
|
||||
"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(),
|
||||
"unsafe_call": UnsafeCallHandler(),
|
||||
"dict": DictHandler(),
|
||||
"comment": CommentHandler(),
|
||||
"set_var": SetVarHandler(),
|
||||
"unsafe_call": UnsafeCallHandler(),
|
||||
}
|
||||
|
||||
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|
||||
|
||||
+98
-30
@@ -3,6 +3,10 @@ class Empty:
|
||||
return False
|
||||
|
||||
|
||||
class ExpReturn(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class ExpBase:
|
||||
def __bool__(self):
|
||||
return True
|
||||
@@ -41,8 +45,10 @@ class ExpSym(str, ExpBase):
|
||||
class ExpTuple(tuple, ExpBase):
|
||||
__slots__ = ()
|
||||
|
||||
def clone(self):
|
||||
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
|
||||
def clone(self, **kwargs):
|
||||
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):
|
||||
val = super().__getitem__(k)
|
||||
@@ -77,11 +83,10 @@ class ExpKV(ExpBase):
|
||||
class ExpDict(dict, ExpBase):
|
||||
__slots__ = ()
|
||||
|
||||
def clone(self):
|
||||
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
|
||||
|
||||
def pop(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
def clone(self, **kwargs):
|
||||
return self.__class__(
|
||||
v.clone(**kwargs) if isinstance(v, ExpBase) else v for v in self
|
||||
)
|
||||
|
||||
def get_eval(self, k, handlers, *args, default=Empty, **kwargs):
|
||||
val = super().get(k, default)
|
||||
@@ -106,12 +111,17 @@ class ExpDict(dict, ExpBase):
|
||||
for k, v in self.items()
|
||||
}
|
||||
|
||||
popitem = pop
|
||||
update = pop
|
||||
clear = pop
|
||||
__delitem__ = pop
|
||||
__setitem__ = pop
|
||||
__ior__ = pop
|
||||
# Can't remember if there was a compelling reason ExpDict can't be mutable but
|
||||
# it breaks deep copy stuff.
|
||||
#
|
||||
# def pop(self, *args, **kwargs):
|
||||
# raise NotImplementedError
|
||||
# popitem = pop
|
||||
# update = pop
|
||||
# clear = pop
|
||||
# __delitem__ = pop
|
||||
# __setitem__ = pop
|
||||
# __ior__ = pop
|
||||
|
||||
|
||||
class ExpStatements(ExpBase):
|
||||
@@ -135,26 +145,81 @@ class ExpStatements(ExpBase):
|
||||
|
||||
|
||||
class ExprGetter:
|
||||
def __init__(self, obj, ctx, *args, **kwargs):
|
||||
def __init__(self, obj, ctx, args, kwargs, *, prepend_args=()):
|
||||
self.obj = obj
|
||||
self.ctx = ctx
|
||||
self.args = args
|
||||
self.prepend_args = prepend_args
|
||||
self.kwargs = kwargs
|
||||
|
||||
def __call__(self, k, *, default=Empty):
|
||||
obj = self.obj
|
||||
result = (
|
||||
obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs)
|
||||
if isinstance(k, str)
|
||||
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
|
||||
)
|
||||
if isinstance(k, str):
|
||||
result = obj.kwargs.get_eval(
|
||||
k, self.ctx, *self.args, default=default, **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:
|
||||
raise KeyError(f"Unknown key {k!r}")
|
||||
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):
|
||||
__slots__ = ("name", "args", "kwargs")
|
||||
__slots__ = ("args", "kwargs", "name")
|
||||
|
||||
def __init__(self, name, args=None, kwargs=None):
|
||||
self.name = name
|
||||
@@ -165,12 +230,14 @@ class ExpFunAp(ExpBase):
|
||||
handler = handlers.get_handler(self.name)
|
||||
if handler is Empty:
|
||||
raise KeyError(f"No handler for op: {self.name!r}")
|
||||
return handler(
|
||||
self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs
|
||||
)
|
||||
return handler(self, getter=ExprGetter(self, handlers, args, kwargs), **kwargs)
|
||||
|
||||
def clone(self):
|
||||
return self.__class__(self.name, self.args.clone(), self.kwargs.clone())
|
||||
def clone(self, **kwargs):
|
||||
return self.__class__(
|
||||
self.name,
|
||||
self.args.clone(**kwargs),
|
||||
self.kwargs.clone(**kwargs),
|
||||
)
|
||||
|
||||
def pretty_string(self, depth=0):
|
||||
pad = " " * (depth + 1) * 2
|
||||
@@ -202,12 +269,13 @@ class ExpBoundFunAp(ExpFunAp):
|
||||
|
||||
__all__ = (
|
||||
"ExpBase",
|
||||
"ExpOp",
|
||||
"ExpBinOp",
|
||||
"ExpSym",
|
||||
"ExpTuple",
|
||||
"ExpKV",
|
||||
"ExpBoundFunAp",
|
||||
"ExpDict",
|
||||
"ExpFunAp",
|
||||
"ExpBoundFunAp",
|
||||
"ExpKV",
|
||||
"ExpMethodAp",
|
||||
"ExpOp",
|
||||
"ExpSym",
|
||||
"ExpTuple",
|
||||
)
|
||||
|
||||
@@ -2,12 +2,12 @@ import contextlib
|
||||
import functools
|
||||
|
||||
from ..latent import ImageBatch
|
||||
from .util import torch
|
||||
from .types import Empty
|
||||
from .util import torch
|
||||
|
||||
|
||||
class Arg:
|
||||
__slots__ = ("name", "default", "validator")
|
||||
__slots__ = ("default", "name", "validator")
|
||||
|
||||
def __init__(self, name, default=Empty, *, validator=None):
|
||||
self.name = name
|
||||
@@ -53,6 +53,23 @@ class Arg:
|
||||
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
|
||||
def tensor_slice(cls, name, default=Empty):
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
def present(cls, name):
|
||||
return cls(name, validator=ValidateArg.validate_passthrough)
|
||||
def present(cls, name, default=Empty):
|
||||
return cls(name, default=default, validator=ValidateArg.validate_passthrough)
|
||||
|
||||
@classmethod
|
||||
def one_of(cls, name, validators, *, default=Empty):
|
||||
@@ -109,7 +126,7 @@ class ValidateError(Exception):
|
||||
|
||||
|
||||
class ValidateArg:
|
||||
__slots__ = ("valfuns", "groupfun", "kwargs", "kwargslist")
|
||||
__slots__ = ("groupfun", "kwargs", "kwargslist", "valfuns")
|
||||
|
||||
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
|
||||
if not isinstance(name, (list, tuple)):
|
||||
@@ -224,6 +241,10 @@ class ValidateArg:
|
||||
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
|
||||
def validate_tensor_slice(cls, idx, val):
|
||||
return cls.validate_sequence(
|
||||
@@ -236,6 +257,12 @@ class ValidateArg:
|
||||
raise ValidateError(f"Expected string argument at {idx}, got {type(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
|
||||
def validate_boolean(cls, idx, val):
|
||||
if val is not True and val is not False:
|
||||
|
||||
+456
-49
@@ -1,18 +1,23 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import PIL.Image as PILImage
|
||||
from functools import partial
|
||||
import torch
|
||||
|
||||
from . import expression as expr
|
||||
from . import latent
|
||||
from . import unsafe_expression_whitelists
|
||||
|
||||
from . import latent, unsafe_expression_whitelists
|
||||
from .external import MODULES as EXT
|
||||
from .utils import scale_noise, resolve_value, quantile_normalize
|
||||
from .latent import OCSTAESD, ImageBatch, normalize_to_scale
|
||||
from .latent import (
|
||||
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_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)
|
||||
|
||||
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor"),
|
||||
@@ -107,6 +126,20 @@ class ClampHandler(NormHandler):
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor),
|
||||
@@ -135,17 +168,66 @@ class ReshapeHandler(NormHandler):
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor_dest"),
|
||||
expr.Arg.tensor("tensor_src"),
|
||||
expr.Arg.tensor_slice("slice"),
|
||||
expr.Arg.boolean("slice_src", default=True),
|
||||
)
|
||||
|
||||
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[tensor_slice] = tensor2[tensor_slice]
|
||||
result[tensor_slice] = tensor2[tensor_slice] if slice_src else tensor2
|
||||
return result
|
||||
|
||||
|
||||
@@ -167,40 +249,36 @@ class NewLikeHandler(NormHandler):
|
||||
class MeanHandler(NormHandler):
|
||||
input_validators = (
|
||||
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):
|
||||
tensor, dim = self.safe_get_all(obj, getter)
|
||||
dim = dim if isinstance(dim, tuple) else (dim,)
|
||||
return tensor.mean(keepdim=True, dim=dim)
|
||||
|
||||
|
||||
class StdHandler(NormHandler):
|
||||
input_validators = (
|
||||
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):
|
||||
tensor, dim = self.safe_get_all(obj, getter)
|
||||
return tensor.std(keepdim=True, dim=dim)
|
||||
tensor, dim, eps = self.safe_get_all(obj, getter)
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor"),
|
||||
expr.Arg.numeric_scalar("amount", 0.5),
|
||||
expr.Arg.one_of(
|
||||
"dim",
|
||||
(
|
||||
expr.ValidateArg.validate_integer,
|
||||
partial(
|
||||
expr.ValidateArg.validate_sequence,
|
||||
item_validator=expr.ValidateArg.validate_integer,
|
||||
),
|
||||
),
|
||||
default=-2,
|
||||
),
|
||||
expr.Arg.numscalar_sequence_or_single("dim", -2),
|
||||
)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
@@ -275,20 +353,166 @@ class NewFullHandler(NormHandler):
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor1"),
|
||||
expr.Arg.tensor("tensor2"),
|
||||
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):
|
||||
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)
|
||||
if not blend_handler:
|
||||
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):
|
||||
@@ -348,7 +572,13 @@ class NoiseHandler(NormHandler):
|
||||
ctx.get_var(k, default=0.0)
|
||||
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)
|
||||
|
||||
|
||||
@@ -360,6 +590,57 @@ class ShapeHandler(expr.BaseHandler):
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor"),
|
||||
@@ -385,6 +666,40 @@ class SNFGuidanceHandler(NormHandler):
|
||||
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):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("reference"),
|
||||
@@ -626,42 +941,134 @@ class ScaleNNLatentUpscaleHandler(expr.BaseHandler):
|
||||
return latent.scale_nnlatentupscale(mode, tensor, scale)
|
||||
|
||||
|
||||
TENSOR_OP_HANDLERS = {
|
||||
"t_norm": NormHandler(),
|
||||
"t_quantilenorm": QuantileNormHandler(),
|
||||
"t_normtoscale": NormToScaleHandler(),
|
||||
"t_normalize_to_scale": NormToScaleHandler(),
|
||||
"t_reshape": ReshapeHandler(),
|
||||
"t_clamp": ClampHandler(),
|
||||
"t_cat": CatHandler(),
|
||||
"t_stack": StackHandler(),
|
||||
"t_indexed_copy": IndexedCopyHandler(),
|
||||
"t_new_like": NewLikeHandler(),
|
||||
"t_mean": MeanHandler(),
|
||||
"t_std": StdHandler(),
|
||||
class ForkRngHandler(expr.BaseHandler):
|
||||
input_validators = (
|
||||
expr.Arg.present("expression"),
|
||||
expr.Arg.one_of(
|
||||
"seed",
|
||||
(
|
||||
expr.ValidateArg.validate_none,
|
||||
expr.ValidateArg.validate_integer,
|
||||
),
|
||||
default=None,
|
||||
),
|
||||
expr.Arg.boolean("enabled", default=True),
|
||||
)
|
||||
|
||||
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_roll": RollHandler(),
|
||||
"t_flip": FlipHandler(),
|
||||
"t_cat": CatHandler(),
|
||||
"t_clamp": ClampHandler(),
|
||||
"t_soft_clamp": SoftClampHandler(),
|
||||
"t_clone": CloneHandler(),
|
||||
"t_newfull": NewFullHandler(),
|
||||
"t_copysign": CopySignHandler(),
|
||||
"t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(),
|
||||
"t_scale": ScaleHandler(),
|
||||
"t_noise": NoiseHandler(),
|
||||
"t_shape": ShapeHandler(),
|
||||
"t_copysign": CopySignHandler(),
|
||||
"t_correlate": CorrelateHandler(),
|
||||
"t_cumsum": CumSumHandler(),
|
||||
"t_flatten": FlattenHandler(),
|
||||
"t_flip": FlipHandler(),
|
||||
"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_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_split": SplitHandler(),
|
||||
"t_stack": StackHandler(),
|
||||
"t_std": StdHandler(),
|
||||
"t_taesd_decode": TAESDDecodeHandler(),
|
||||
"t_trim": TrimHandler(),
|
||||
"unsafe_tensor_method": UnsafeTorchTensorMethodHandler(),
|
||||
"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 = {
|
||||
"img_taesd_encode": TAESDEncodeHandler(),
|
||||
"img_shape": ImgShapeHandler(),
|
||||
"img_pil_resize": ImgPILResizeHandler(),
|
||||
}
|
||||
|
||||
HANDLERS |= TENSOR_OP_HANDLERS
|
||||
HANDLERS |= TORCH_OP_HANDLERS
|
||||
HANDLERS |= IMAGE_OP_HANDLERS
|
||||
|
||||
+316
-9
@@ -1,13 +1,13 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from typing import Any, NamedTuple, Self
|
||||
|
||||
import folder_paths
|
||||
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.utils import bislerp
|
||||
from comfy import latent_formats
|
||||
|
||||
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
|
||||
# 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.
|
||||
def contrast_adaptive_sharpening( # noqa: PLR0914
|
||||
def contrast_adaptive_sharpening(
|
||||
x,
|
||||
amount=0.8,
|
||||
*,
|
||||
@@ -228,7 +228,7 @@ def scale_samples(
|
||||
|
||||
def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): # noqa: F811
|
||||
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)
|
||||
|
||||
|
||||
@@ -324,7 +324,7 @@ class OCSLatentFormat:
|
||||
|
||||
def latent_to_rgb(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
# NCHW -> NHWC
|
||||
if self.latent_factors is None:
|
||||
if self.rgb_factors is None:
|
||||
raise ValueError("No RGB factors for latent type!")
|
||||
return torch.nn.functional.linear(
|
||||
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:
|
||||
# NHWC
|
||||
if self.latent_factors is None:
|
||||
if self.rgb_factors is None:
|
||||
raise ValueError("No RGB factors for latent type!")
|
||||
if self.rgb_factors_bias is not None:
|
||||
img = img - self.rgb_factors_bias
|
||||
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.cfg1_uncond_optimization = cfg1_uncond_optimization
|
||||
self.cfg_scale_override = cfg_scale_override
|
||||
self.model_sampling = model.inner_model.inner_model.model_sampling
|
||||
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(
|
||||
x.device, model.inner_model.inner_model.latent_format
|
||||
@@ -212,9 +213,9 @@ class OCSModel:
|
||||
) -> torch.Tensor:
|
||||
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
||||
|
||||
@property
|
||||
def model_sampling(self):
|
||||
return self.model.inner_model.inner_model.model_sampling
|
||||
# @property
|
||||
# def model_sampling(self):
|
||||
# return self.model.inner_model.inner_model.model_sampling
|
||||
|
||||
@property
|
||||
def inner_cfg_scale(self) -> None | int | float:
|
||||
|
||||
+343
-13
@@ -1,17 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import comfy
|
||||
import yaml
|
||||
|
||||
import torch
|
||||
|
||||
import yaml
|
||||
from tqdm import tqdm
|
||||
|
||||
from .external import MODULES, IntegratedNode
|
||||
from .filtering import Filter, FilterRefs, make_filter
|
||||
from .restart import Restart
|
||||
from .sampling import composable_sampler
|
||||
from .step_samplers import STEP_SAMPLERS
|
||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||
from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups
|
||||
from .filtering import make_filter
|
||||
|
||||
try:
|
||||
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}"
|
||||
)
|
||||
|
||||
PARAM_INPUT_TYPES = frozenset((
|
||||
"IMAGE",
|
||||
"OCS_NOISE",
|
||||
"SAMPLER",
|
||||
"SIGMAS",
|
||||
"SONAR_CUSTOM_NOISE",
|
||||
"UPSCALE_MODEL",
|
||||
"VAE",
|
||||
))
|
||||
PARAM_INPUT_TYPES = frozenset(
|
||||
(
|
||||
"IMAGE",
|
||||
"OCS_NOISE",
|
||||
"SAMPLER",
|
||||
"SIGMAS",
|
||||
"SONAR_CUSTOM_NOISE",
|
||||
"UPSCALE_MODEL",
|
||||
"VAE",
|
||||
)
|
||||
)
|
||||
|
||||
NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
|
||||
|
||||
@@ -778,6 +780,334 @@ class ApplyFilterImage(ApplyFilterLatent):
|
||||
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__ = (
|
||||
"SamplerNode",
|
||||
"GroupNode",
|
||||
|
||||
+84
-19
@@ -2,12 +2,67 @@ import gc
|
||||
import math
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import scipy
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
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):
|
||||
@@ -19,6 +74,12 @@ class ImmiscibleNoise(Filter):
|
||||
"maximize": False,
|
||||
"distance_scale": 0.0,
|
||||
"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):
|
||||
@@ -83,7 +144,11 @@ class ImmiscibleNoise(Filter):
|
||||
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
|
||||
# Minimize latent-noise pairs over a batch
|
||||
batch = latent.shape[0]
|
||||
out_latent = fallback(out_latent, latent)
|
||||
ref_latent = ref_latent.detach().clone()
|
||||
if self.abs_mode:
|
||||
ref_latent = ref_latent.abs()
|
||||
latent = latent.abs()
|
||||
if self.distance_scale == 0:
|
||||
ref_latent_expanded = ref_latent.unsqueeze(1).expand(
|
||||
-1, batch, *ref_latent.shape[1:]
|
||||
@@ -92,30 +157,30 @@ class ImmiscibleNoise(Filter):
|
||||
ref_latent.shape[0], *latent.shape
|
||||
)
|
||||
dist = (ref_latent_expanded - latent_expanded) ** 2
|
||||
if self.abs_distance_mode:
|
||||
dist = dist.abs_()
|
||||
del ref_latent_expanded, latent_expanded
|
||||
dist = dist.mean(tuple(range(2, dist.dim())))
|
||||
dist = dist.mean(tuple(range(2, dist.ndim)))
|
||||
else:
|
||||
dist = torch.linalg.vector_norm(
|
||||
fallback(self.distance_scale_ref, self.distance_scale)
|
||||
* ref_latent.flatten(start_dim=1).unsqueeze(1)
|
||||
- self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0),
|
||||
dim=2,
|
||||
)
|
||||
dist = dist.half()
|
||||
distance_scale_ref = fallback(self.distance_scale_ref, self.distance_scale)
|
||||
dist = distance_scale_ref * ref_latent.flatten(start_dim=1).unsqueeze(
|
||||
1
|
||||
) - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0)
|
||||
if self.abs_distance_mode:
|
||||
dist = dist.abs_()
|
||||
dist = torch.linalg.vector_norm(dist, dim=2)
|
||||
try:
|
||||
assign_mat = scipy.optimize.linear_sum_assignment(
|
||||
dist.cpu(), maximize=self.maximize
|
||||
assign_mat = linear_sum_assignment(
|
||||
dist,
|
||||
maximize=self.maximize,
|
||||
use_triton=self.use_triton,
|
||||
split_batch=self.split_batch,
|
||||
generator=self.generator,
|
||||
)
|
||||
except ValueError as exc:
|
||||
tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}")
|
||||
return (
|
||||
None
|
||||
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]]
|
||||
)
|
||||
return None if return_idxs else out_latent[: ref_latent.shape[0]]
|
||||
return assign_mat if return_idxs else out_latent[assign_mat[1]]
|
||||
|
||||
def immiscible_simple(
|
||||
self,
|
||||
|
||||
+112
-10
@@ -1,8 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import NamedTuple
|
||||
|
||||
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:
|
||||
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
|
||||
|
||||
self.s_noise = s_noise
|
||||
@@ -10,6 +59,9 @@ class Restart:
|
||||
immiscible = ImmiscibleNoise(**immiscible)
|
||||
self.immiscible = immiscible
|
||||
self.custom_noise = custom_noise
|
||||
self.normalized = normalized
|
||||
self.normalize_dims = normalize_dims
|
||||
self.is_flow = is_flow
|
||||
|
||||
def get_noise_sampler(self, nsc):
|
||||
return nsc.make_caching_noise_sampler(
|
||||
@@ -30,23 +82,64 @@ class Restart:
|
||||
last_sigma = sigma
|
||||
return sigmas
|
||||
|
||||
def split_sigmas(self, sigmas):
|
||||
def split_sigmas(self, sigmas: torch.Tensor):
|
||||
prev_seg = None
|
||||
while len(sigmas) > 1:
|
||||
seg = self.get_segment(sigmas)
|
||||
sigmas = sigmas[len(seg) :]
|
||||
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:
|
||||
noise_scale = 0.0
|
||||
scale_factors = None
|
||||
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
|
||||
if isinstance(result, torch.Tensor):
|
||||
result = result.item()
|
||||
return result * self.s_noise
|
||||
return result.item()
|
||||
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):
|
||||
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")
|
||||
sched_idx = item
|
||||
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:
|
||||
break
|
||||
interval, jump = item
|
||||
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}")
|
||||
out += chunk
|
||||
sig_idx += interval + jump
|
||||
if jump >= 0:
|
||||
sig_idx += 1
|
||||
sig_idx += interval + jump
|
||||
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:
|
||||
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]:
|
||||
out.append(siglist[-1])
|
||||
return torch.tensor(out).to(sigmas)
|
||||
|
||||
+31
-21
@@ -1,13 +1,12 @@
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
|
||||
|
||||
from .filtering import FILTER_HANDLERS, FilterRefs
|
||||
from .model import OCSModel
|
||||
from .noise import NoiseSamplerCache
|
||||
from .substep_sampling import SamplerState
|
||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||
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:
|
||||
@@ -48,11 +47,6 @@ def composable_sampler(
|
||||
restart_custom_noise = copts.get("restart_custom_noise")
|
||||
if isinstance(restart_custom_noise, str):
|
||||
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(
|
||||
OCSModel(
|
||||
@@ -72,6 +66,16 @@ def composable_sampler(
|
||||
reta=copts.get("reta", 1.0),
|
||||
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"]
|
||||
merge_samplers = tuple(
|
||||
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
|
||||
@@ -85,17 +89,17 @@ def composable_sampler(
|
||||
)
|
||||
ss.noise = nsc
|
||||
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)
|
||||
ss.total_steps = step_count
|
||||
step = 0
|
||||
with trange(step_count, disable=ss.disable_status) as pbar:
|
||||
for noise_scale, chunk_sigmas in sigma_chunks:
|
||||
if step != 0 and noise_scale != 0:
|
||||
prev_refs = FilterRefs({
|
||||
f"pre_restart_{k}": v for k, v in ss.refs.items()
|
||||
})
|
||||
for chunk_idx, (scale_factors, chunk_sigmas) in enumerate(sigma_chunks):
|
||||
if step != 0 and scale_factors is not None:
|
||||
prev_refs = FilterRefs(
|
||||
{f"pre_restart_{k}": v for k, v in ss.refs.items()}
|
||||
)
|
||||
ss.sigmas = chunk_sigmas
|
||||
ss.update(0, step=step, substep=0)
|
||||
if step != 0:
|
||||
@@ -104,14 +108,20 @@ def composable_sampler(
|
||||
ss.hist.reset()
|
||||
for ms in merge_samplers:
|
||||
ms.reset()
|
||||
nsc.min_sigma, nsc.max_sigma = chunk_sigmas[-1], chunk_sigmas[0]
|
||||
if step != 0 and noise_scale != 0:
|
||||
restart_ns = restart.get_noise_sampler(nsc)
|
||||
x += nsc.scale_noise(
|
||||
restart_ns(nsc.min_sigma, nsc.max_sigma, refs=prev_refs | ss.refs),
|
||||
noise_scale,
|
||||
nsc.min_sigma, nsc.max_sigma = (
|
||||
chunk_sigmas[-1].clone(),
|
||||
chunk_sigmas[0].clone(),
|
||||
)
|
||||
if step != 0 and scale_factors is not None:
|
||||
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
|
||||
for idx in range(len(chunk_sigmas) - 1):
|
||||
if idx > 0:
|
||||
|
||||
@@ -1,20 +1,19 @@
|
||||
import inspect
|
||||
import math
|
||||
import typing
|
||||
|
||||
import inspect
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
import comfy
|
||||
|
||||
from .. import filtering
|
||||
from .. import expression as expr
|
||||
from .. import filtering
|
||||
from ..utils import fallback
|
||||
from .base import (
|
||||
StepSamplerContext,
|
||||
SingleStepSampler,
|
||||
StepSamplerContext,
|
||||
registry,
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
import pytorch_wavelets as ptwav
|
||||
|
||||
@@ -531,7 +530,10 @@ class WeoonStep(SingleStepSampler):
|
||||
)
|
||||
denoised_new = self.wavelet_inverse(coeffs_out)
|
||||
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)
|
||||
yield from self.result(x, sigma_up)
|
||||
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
import torch
|
||||
|
||||
import comfy
|
||||
import torch
|
||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||
from tqdm import tqdm
|
||||
|
||||
from .. import filtering
|
||||
from .base import (
|
||||
SingleStepSampler,
|
||||
DPMPPStepMixin,
|
||||
HistorySingleStepSampler,
|
||||
ReversibleSingleStepSampler,
|
||||
SingleStepSampler,
|
||||
registry,
|
||||
)
|
||||
|
||||
@@ -524,6 +525,147 @@ class RESMultistepStep(HistorySingleStepSampler, DPMPPStepMixin):
|
||||
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(
|
||||
DEISStep,
|
||||
DPMPP2MSDEStep,
|
||||
@@ -538,4 +680,5 @@ registry.add(
|
||||
DPM2Step,
|
||||
DPMPP2SStep,
|
||||
RESMultistepStep,
|
||||
Seeds2Step,
|
||||
)
|
||||
|
||||
@@ -10,7 +10,7 @@ class PingPongStep(SingleStepSampler):
|
||||
super().__init__(*args, **kwargs)
|
||||
pingpong_options = self.options.pop("pingpong", {})
|
||||
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):
|
||||
ss = self.ss
|
||||
|
||||
+16
-11
@@ -5,8 +5,7 @@ import tqdm
|
||||
|
||||
from . import expression as expr
|
||||
from . import utils
|
||||
|
||||
from .filtering import make_filter, FilterRefs, FILTER_HANDLERS
|
||||
from .filtering import FILTER_HANDLERS, FilterRefs, make_filter
|
||||
from .noise import ImmiscibleNoise
|
||||
from .restart import Restart
|
||||
from .step_samplers import STEP_SAMPLERS
|
||||
@@ -451,6 +450,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
s_noise=restart.get("s_noise", 1.0),
|
||||
custom_noise=restart_custom_noise,
|
||||
immiscible=restart.get("immiscible", False),
|
||||
is_flow=ss.model.is_rectified_flow,
|
||||
)
|
||||
|
||||
def make_schedule(self, ss):
|
||||
@@ -505,10 +505,13 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
if subss.idx >= max_idx:
|
||||
break
|
||||
if last_down is not None and last_down < ss.sigma_next:
|
||||
restart_ns = self.restart.get_noise_sampler(ss.noise)
|
||||
x += ss.noise.scale_noise(
|
||||
restart_ns(last_down, ss.sigma_next, refs=ss.refs),
|
||||
self.restart.get_noise_scale(last_down, ss.sigma_next),
|
||||
x = self.restart.add_noise(
|
||||
x,
|
||||
sigma_from=last_down.item(),
|
||||
sigma_to=ss.sigma_next.item(),
|
||||
nsc=nsc,
|
||||
refs=ss.refs,
|
||||
in_place=True,
|
||||
)
|
||||
pbar.update(0)
|
||||
return x
|
||||
@@ -653,11 +656,13 @@ class PingpongMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
sigma_next,
|
||||
immiscible=fallback(self.immiscible, ss.noise.immiscible),
|
||||
)
|
||||
noise_refs = ss.refs | FilterRefs({
|
||||
"orig_x": orig_x,
|
||||
"x": x,
|
||||
"denoised": synth_denoised,
|
||||
})
|
||||
noise_refs = ss.refs | FilterRefs(
|
||||
{
|
||||
"orig_x": orig_x,
|
||||
"x": x,
|
||||
"denoised": synth_denoised,
|
||||
}
|
||||
)
|
||||
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 .filtering import FilterRefs
|
||||
@@ -7,6 +8,13 @@ from .model import History
|
||||
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:
|
||||
def __init__(self, items=None):
|
||||
self.items = [] if items is None else items
|
||||
@@ -141,6 +149,10 @@ class SamplerState:
|
||||
self.substep = 0
|
||||
self.total_steps = len(sigmas) - 1
|
||||
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
|
||||
|
||||
@property
|
||||
@@ -171,6 +183,23 @@ class SamplerState:
|
||||
def d(self):
|
||||
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):
|
||||
idx = self.idx if idx is None else idx
|
||||
self.idx = idx
|
||||
@@ -185,6 +214,59 @@ class SamplerState:
|
||||
self.substep = substep
|
||||
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(
|
||||
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
|
||||
else:
|
||||
preview = hi.denoised
|
||||
return self.callback_({
|
||||
"x": hi.x,
|
||||
"i": self.step,
|
||||
"sigma": hi.sigma,
|
||||
"sigma_hat": hi.sigma,
|
||||
"denoised": preview,
|
||||
})
|
||||
return self.callback_(
|
||||
{
|
||||
"x": hi.x,
|
||||
"i": self.step,
|
||||
"sigma": hi.sigma,
|
||||
"sigma_hat": hi.sigma,
|
||||
"denoised": preview,
|
||||
}
|
||||
)
|
||||
|
||||
def reset(self):
|
||||
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 torch
|
||||
|
||||
from comfy.k_diffusion.sampling import to_d
|
||||
|
||||
# def scale_noise_(
|
||||
# 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)
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
def scale_noise(
|
||||
noise,
|
||||
factor=1.0,
|
||||
noise: torch.Tensor,
|
||||
factor: float = 1.0,
|
||||
*,
|
||||
normalized=True,
|
||||
normalize_dims=(-3, -2, -1),
|
||||
):
|
||||
normalized: bool = True,
|
||||
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:
|
||||
return noise * factor if factor != 1 else noise
|
||||
noise = noise / noise.std(dim=normalize_dims, keepdim=True)
|
||||
return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor)
|
||||
|
||||
|
||||
def _quantile_norm_scaledown(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
**_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",
|
||||
if eps is None:
|
||||
eps = torch.finfo(noise.dtype).eps * 1.25
|
||||
if normalize_dims is None:
|
||||
normalize_dims = tuple(
|
||||
range(
|
||||
max(0, min(1, noise.ndim - 1)),
|
||||
noise.ndim,
|
||||
)
|
||||
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(
|
||||
# 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
|
||||
# n = (
|
||||
# torch.nn.LayerNorm(noise.shape[1:])
|
||||
# if normalize_dims == (-3, -2, -1)
|
||||
# else torch.nn.InstanceNorm2d(noise.shape[1])
|
||||
# ).to(noise)
|
||||
# return n(noise) * factor
|
||||
# return latent.normalize_to_scale(
|
||||
# n(noise).clamp_(-1, 1), -1, 1, dim=normalize_dims
|
||||
# ).mul_(factor)
|
||||
def range_wrap(
|
||||
x: torch.Tensor,
|
||||
min_val: float | torch.Tensor,
|
||||
max_val: float | torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return min_val + (x - min_val).remainder_(max_val - min_val)
|
||||
|
||||
|
||||
def softplus_soft_clamp(
|
||||
t: torch.Tensor,
|
||||
min_val: torch.Tensor | float = 0.0,
|
||||
max_val: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
# We define stiffness as a multiplier (beta) for the softplus function.
|
||||
# Higher stiffness = sharper transition.
|
||||
stiffness: float = 1.0,
|
||||
safe: bool = True,
|
||||
) -> 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):
|
||||
|
||||
@@ -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