3 Commits
Author SHA1 Message Date
blepping 79c7b24d1e Expand expression handlers, more functionality for ExpressionFiltered nodes 2026-09-08 03:28:17 -06:00
blepping 5ea059bed5 Bug fixes
More expression tensor operations
Make the return expression handler actually work
2026-07-11 10:44:15 -06:00
blepping ee59df94e3 New Perlin, expression filtered patches, expression methods 2026-06-23 11:15:31 -06:00
24 changed files with 3794 additions and 1055 deletions
+2
View File
@@ -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"]
+3 -3
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+8 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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",
)
+32 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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:
+9 -7
View File
@@ -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)
+146 -3
View File
@@ -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,
)
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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()
+386
View File
@@ -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
View File
@@ -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):
+238
View File
@@ -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])),
)