New Perlin, expression filtered patches, expression methods
This commit is contained in:
@@ -11,5 +11,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
|
||||
"OCS ApplyFilterLatent": nodes.ApplyFilterLatent,
|
||||
"OCS ApplyFilterImage": nodes.ApplyFilterImage,
|
||||
"OCS ExpressionFilteredLatentOperation": nodes.ExpressionFilteredLatentOperationNode,
|
||||
"OCS ExpressionFilteredModelPatch": nodes.ExpressionFilteredModelPatchNode,
|
||||
} | custom_noise.NODE_CLASS_MAPPINGS
|
||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import abc
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from typing import Callable, Any
|
||||
|
||||
from ..external import IntegratedNode
|
||||
from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE
|
||||
from ..noise import scale_noise
|
||||
from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT
|
||||
|
||||
|
||||
class CustomNoiseItemBase(abc.ABC):
|
||||
|
||||
+552
-236
File diff suppressed because it is too large
Load Diff
+578
-419
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,12 @@
|
||||
from . import types, expression, handler, util, validation
|
||||
|
||||
from . import expression, types, util
|
||||
from .expression import Expression
|
||||
from .validation import Arg, ValidateArg
|
||||
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
|
||||
|
||||
try:
|
||||
from . import handler, validation
|
||||
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
|
||||
from .validation import Arg, ValidateArg
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
pass
|
||||
|
||||
__all__ = (
|
||||
"Arg",
|
||||
|
||||
+22
-11
@@ -1,20 +1,21 @@
|
||||
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,
|
||||
ExpStatements,
|
||||
ExpSym,
|
||||
ExpTuple,
|
||||
)
|
||||
|
||||
COMMA_PRECEDENCE = 2
|
||||
@@ -36,10 +37,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*
|
||||
@@ -157,9 +159,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 +202,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 +261,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, ("*", "/"))
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import operator
|
||||
import traceback
|
||||
|
||||
from .validation import ValidateArg, Arg, ValidateError
|
||||
from .types import Empty, ExpDict, ExpOp
|
||||
from .types import Empty, ExpDict, ExpOp, ExpReturn
|
||||
from .util import torch
|
||||
from .validation import Arg, ValidateArg, ValidateError
|
||||
|
||||
|
||||
class HandlerError(Exception):
|
||||
@@ -63,7 +64,8 @@ class BaseHandler:
|
||||
val = self.handle(obj, getter)
|
||||
return self.validate_output(obj, val)
|
||||
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
|
||||
|
||||
def safe_get(self, key, obj, getter=None, *, default=Empty):
|
||||
str_key = isinstance(key, str)
|
||||
@@ -365,6 +367,13 @@ 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))
|
||||
|
||||
|
||||
LOGIC_HANDLERS = {
|
||||
"||": OrHandler(),
|
||||
"&&": AndHandler(),
|
||||
@@ -400,6 +409,9 @@ MATH_HANDLERS = {
|
||||
">=": RelComparisonHandler(operator.ge),
|
||||
"min": MinHandler(),
|
||||
"max": MaxHandler(),
|
||||
"float": UnarySimpleMathHandler(handler=float),
|
||||
"int": UnarySimpleMathHandler(handler=int),
|
||||
"bool": UnarySimpleMathHandler(handler=bool),
|
||||
}
|
||||
for k, alias in (
|
||||
("+", "add"),
|
||||
@@ -420,6 +432,7 @@ MISC_HANDLERS = {
|
||||
"dict": DictHandler(),
|
||||
"comment": CommentHandler(),
|
||||
"set_var": SetVarHandler(),
|
||||
"return": ReturnHandler(),
|
||||
}
|
||||
|
||||
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|
||||
|
||||
+75
-17
@@ -2,6 +2,8 @@ class Empty:
|
||||
def __bool__(self):
|
||||
return False
|
||||
|
||||
class ExpReturn(Exception):
|
||||
pass
|
||||
|
||||
class ExpBase:
|
||||
def __bool__(self):
|
||||
@@ -80,8 +82,6 @@ class ExpDict(dict, ExpBase):
|
||||
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 get_eval(self, k, handlers, *args, default=Empty, **kwargs):
|
||||
val = super().get(k, default)
|
||||
@@ -106,12 +106,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,24 +140,78 @@ 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__ = ("object_expression", "funap")
|
||||
|
||||
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):
|
||||
return self.__class__(
|
||||
object_expression=self.object_expression.clone(),
|
||||
funap=self.funap.clone(),
|
||||
)
|
||||
__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")
|
||||
|
||||
@@ -165,9 +224,7 @@ 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())
|
||||
@@ -209,5 +266,6 @@ __all__ = (
|
||||
"ExpKV",
|
||||
"ExpDict",
|
||||
"ExpFunAp",
|
||||
"ExpMethodAp",
|
||||
"ExpBoundFunAp",
|
||||
)
|
||||
|
||||
@@ -2,8 +2,8 @@ import contextlib
|
||||
import functools
|
||||
|
||||
from ..latent import ImageBatch
|
||||
from .util import torch
|
||||
from .types import Empty
|
||||
from .util import torch
|
||||
|
||||
|
||||
class Arg:
|
||||
@@ -53,6 +53,17 @@ class Arg:
|
||||
name, default=default, validator=ValidateArg.validate_numscalar_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)
|
||||
@@ -236,6 +247,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:
|
||||
|
||||
+86
-27
@@ -1,18 +1,15 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import PIL.Image as PILImage
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image as PILImage
|
||||
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 .utils import quantile_normalize, resolve_value, scale_noise
|
||||
|
||||
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
|
||||
@@ -107,6 +104,14 @@ class ClampHandler(NormHandler):
|
||||
return torch.clamp(tensor, min=tmin, max=tmax)
|
||||
|
||||
|
||||
class AbsHandler(NormHandler):
|
||||
input_validators = (expr.Arg.tensor("tensor"),)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
(tensor,) = self.safe_get_all(obj, getter)
|
||||
return tensor.abs()
|
||||
|
||||
|
||||
class StackHandler(NormHandler):
|
||||
input_validators = (
|
||||
expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor),
|
||||
@@ -135,17 +140,52 @@ 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"),
|
||||
)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
tensor, shape = self.safe_get_all(obj, getter)
|
||||
return tensor[tuple(slice(0, dsize) for dsize in shape)]
|
||||
|
||||
|
||||
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,22 +207,24 @@ 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)),
|
||||
)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
tensor, dim = self.safe_get_all(obj, getter)
|
||||
dim = dim if isinstance(dim, tuple) else (dim,)
|
||||
return tensor.std(keepdim=True, dim=dim)
|
||||
|
||||
|
||||
@@ -190,17 +232,7 @@ 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):
|
||||
@@ -280,15 +312,34 @@ class BlendHandler(NormHandler):
|
||||
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,
|
||||
),
|
||||
),
|
||||
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):
|
||||
@@ -633,8 +684,11 @@ TENSOR_OP_HANDLERS = {
|
||||
"t_normalize_to_scale": NormToScaleHandler(),
|
||||
"t_reshape": ReshapeHandler(),
|
||||
"t_clamp": ClampHandler(),
|
||||
"t_abs": AbsHandler(),
|
||||
"t_cat": CatHandler(),
|
||||
"t_stack": StackHandler(),
|
||||
"t_split": SplitHandler(),
|
||||
"t_trim": TrimHandler(),
|
||||
"t_indexed_copy": IndexedCopyHandler(),
|
||||
"t_new_like": NewLikeHandler(),
|
||||
"t_mean": MeanHandler(),
|
||||
@@ -657,6 +711,11 @@ TENSOR_OP_HANDLERS = {
|
||||
"unsafe_torch": UnsafeTorchHandler(),
|
||||
}
|
||||
|
||||
TENSOR_OP_HANDLERS |= {
|
||||
f"Tensor::{k[2:] if k.startswith('t_') else k}": v
|
||||
for k, v in TENSOR_OP_HANDLERS.items()
|
||||
}
|
||||
|
||||
IMAGE_OP_HANDLERS = {
|
||||
"img_taesd_encode": TAESDEncodeHandler(),
|
||||
"img_shape": ImgShapeHandler(),
|
||||
|
||||
+5
-4
@@ -169,8 +169,9 @@ class OCSModel:
|
||||
self.extra_args = extra_args
|
||||
self.cfg1_uncond_optimization = cfg1_uncond_optimization
|
||||
self.cfg_scale_override = cfg_scale_override
|
||||
self.model_sampling = model.inner_model.inner_model.model_sampling
|
||||
self.is_rectified_flow = isinstance(
|
||||
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
|
||||
self.model_sampling, comfy.model_sampling.CONST
|
||||
)
|
||||
self.latent_format = OCSLatentFormat(
|
||||
x.device, model.inner_model.inner_model.latent_format
|
||||
@@ -212,9 +213,9 @@ class OCSModel:
|
||||
) -> torch.Tensor:
|
||||
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
||||
|
||||
@property
|
||||
def model_sampling(self):
|
||||
return self.model.inner_model.inner_model.model_sampling
|
||||
# @property
|
||||
# def model_sampling(self):
|
||||
# return self.model.inner_model.inner_model.model_sampling
|
||||
|
||||
@property
|
||||
def inner_cfg_scale(self) -> None | int | float:
|
||||
|
||||
+338
-13
@@ -1,17 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import comfy
|
||||
import yaml
|
||||
|
||||
import torch
|
||||
|
||||
import yaml
|
||||
from tqdm import tqdm
|
||||
|
||||
from .external import MODULES, IntegratedNode
|
||||
from .filtering import Filter, FilterRefs, make_filter
|
||||
from .restart import Restart
|
||||
from .sampling import composable_sampler
|
||||
from .step_samplers import STEP_SAMPLERS
|
||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||
from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups
|
||||
from .filtering import make_filter
|
||||
|
||||
try:
|
||||
from comfy_execution import validation as comfy_validation
|
||||
@@ -27,15 +27,17 @@ except Exception as exc:
|
||||
f"** OCS: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}"
|
||||
)
|
||||
|
||||
PARAM_INPUT_TYPES = frozenset((
|
||||
"IMAGE",
|
||||
"OCS_NOISE",
|
||||
"SAMPLER",
|
||||
"SIGMAS",
|
||||
"SONAR_CUSTOM_NOISE",
|
||||
"UPSCALE_MODEL",
|
||||
"VAE",
|
||||
))
|
||||
PARAM_INPUT_TYPES = frozenset(
|
||||
(
|
||||
"IMAGE",
|
||||
"OCS_NOISE",
|
||||
"SAMPLER",
|
||||
"SIGMAS",
|
||||
"SONAR_CUSTOM_NOISE",
|
||||
"UPSCALE_MODEL",
|
||||
"VAE",
|
||||
)
|
||||
)
|
||||
|
||||
NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
|
||||
|
||||
@@ -778,6 +780,329 @@ 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={
|
||||
"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 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 ValueError("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 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 ValueError("Bad type for filter definition, must be object")
|
||||
latent_refs = {
|
||||
k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True)
|
||||
for k, v in (
|
||||
("latent_ref_1", latent_ref_1_opt),
|
||||
("latent_ref_2", latent_ref_2_opt),
|
||||
("latent_ref_3", latent_ref_3_opt),
|
||||
)
|
||||
if v is not None
|
||||
}
|
||||
ocs_filter = make_filter(filter_def)
|
||||
model = model.clone()
|
||||
mode_keys = {
|
||||
"post_cfg": "sampler_post_cfg_function",
|
||||
"pre_cfg": "sampler_pre_cfg_function",
|
||||
"apply_model": "model_function_wrapper",
|
||||
"cfg": "sampler_cfg_function",
|
||||
"denoise_mask": "denoise_mask_function",
|
||||
}
|
||||
key = mode_keys.get(patch_mode)
|
||||
if key is None:
|
||||
raise ValueError(f"Bad mode: {patch_mode}")
|
||||
if existing_patch_mode != "normal":
|
||||
old_handlers = model.model_options.pop(key, None)
|
||||
if old_handlers is None:
|
||||
old_handlers = ()
|
||||
else:
|
||||
old_handlers = ()
|
||||
|
||||
def get_refs(*args, **kwargs) -> FilterRefs:
|
||||
old_results = []
|
||||
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
|
||||
argdict = args[0]
|
||||
elif patch_mode == "denoise_mask":
|
||||
argdict = {
|
||||
"sigma": args[0],
|
||||
"denoise_mask": args[1].clone(),
|
||||
"sigmas": kwargs["extra_options"]["sigmas"].clone(),
|
||||
}
|
||||
elif patch_mode == "apply_model":
|
||||
argdict = args[1] | {"apply_function": args[0]}
|
||||
else:
|
||||
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||
if old_handlers:
|
||||
ridx = 0 if existing_patch_mode == "extract_sequence" else -1
|
||||
if patch_mode in {"cfg", "denoise_mask", "apply_model"}:
|
||||
old_results = (old_handlers[0](*args, **kwargs),)
|
||||
elif patch_mode == "pre_cfg":
|
||||
old_results = [argdict["conds_out"]]
|
||||
for hf in old_handlers:
|
||||
result = hf(argdict | {"conds_out": old_results[ridx]}).copy()
|
||||
if len(old_results) > 1 and existing_patch_mode != "extract":
|
||||
old_results[1] = result
|
||||
else:
|
||||
old_results.append(result)
|
||||
old_results = old_results[1:]
|
||||
elif patch_mode == "post_cfg":
|
||||
old_results = [argdict["denoised"].clone()]
|
||||
for hf in old_handlers:
|
||||
result = hf(argdict | {"denoised": old_results[ridx].clone()})
|
||||
if len(old_results) > 1 and existing_patch_mode != "extract":
|
||||
old_results[1] = result
|
||||
else:
|
||||
old_results.append(result)
|
||||
old_results = old_results[1:]
|
||||
kvs = {
|
||||
"sigma": argdict["sigma"].clone(),
|
||||
"sigma_float": argdict["sigma"].max().item(),
|
||||
"old_results": tuple(old_results),
|
||||
}
|
||||
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
|
||||
kvs |= {
|
||||
"x": argdict["input"].clone(),
|
||||
"cfg_scale": argdict["cond_scale"],
|
||||
}
|
||||
if patch_mode in {"post_cfg", "cfg"}:
|
||||
kvs["cond"] = argdict["cond_denoised"].clone()
|
||||
uncond = argdict.get("uncond_denoised", None)
|
||||
kvs["uncond"] = uncond if uncond is None else uncond.clone()
|
||||
if patch_mode == "post_cfg":
|
||||
kvs["denoised"] = argdict["denoised"].clone()
|
||||
else:
|
||||
conds_out = argdict["conds_out"]
|
||||
kvs["cond"] = conds_out[0].clone()
|
||||
kvs["uncond"] = (
|
||||
conds_out[1].clone()
|
||||
if len(conds_out) > 1 and conds_out[1] is not None
|
||||
else None
|
||||
)
|
||||
kvs["conds_out"] = list(conds_out)
|
||||
elif patch_mode == "denoise_mask":
|
||||
kvs |= {
|
||||
"sigmas": argdict["sigmas"],
|
||||
"denoise_mask": argdict["denoise_mask"],
|
||||
}
|
||||
elif patch_mode == "apply_model":
|
||||
kvs |= {
|
||||
"x": argdict["input"].clone(),
|
||||
"cond_or_uncond": argdict["cond_or_uncond"].clone(),
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
|
||||
latent_in = kvs["x"]
|
||||
else:
|
||||
latent_in = kvs["denoise_mask"]
|
||||
kvs |= {k: v.to(latent_in) for k, v in latent_refs.items()}
|
||||
return FilterRefs(kvs=kvs)
|
||||
|
||||
def model_patch(*args, **kwargs):
|
||||
refs = get_refs(*args, **kwargs)
|
||||
if patch_mode == "apply_model":
|
||||
|
||||
def fallback_apply_model():
|
||||
old_results = refs.kvs["old_results"]
|
||||
if old_results:
|
||||
return old_results[-1]
|
||||
return args[0](
|
||||
args[1]["input"], args[1]["timestep"], **args[1]["c"]
|
||||
)
|
||||
else:
|
||||
fallback_apply_model = None
|
||||
if not ocs_filter.check_applies(refs):
|
||||
old_results = refs.kvs["old_results"]
|
||||
if old_results:
|
||||
return old_results[-1]
|
||||
if patch_mode == "pre_cfg":
|
||||
return args[0]["conds_out"]
|
||||
if patch_mode == "post_cfg":
|
||||
return args[0]["denoised"]
|
||||
if patch_mode == "cfg":
|
||||
return args[0]["cond"]
|
||||
if patch_mode == "denoise_mask":
|
||||
return args[1]
|
||||
if patch_mode == "apply_model":
|
||||
return fallback_apply_model()
|
||||
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
|
||||
latent_in = refs.kvs["x"]
|
||||
else:
|
||||
latent_in = refs.kvs["denoise_mask"]
|
||||
result = ocs_filter.apply(latent_in, refs=refs)
|
||||
if patch_mode == "apply_model" and result is None:
|
||||
return fallback_apply_model()
|
||||
if patch_mode == "pre_cfg":
|
||||
return list(result)
|
||||
return result
|
||||
|
||||
if patch_mode == "pre_cfg":
|
||||
model.set_model_sampler_pre_cfg_function(model_patch)
|
||||
elif patch_mode == "post_cfg":
|
||||
model.set_model_sampler_post_cfg_function(model_patch)
|
||||
elif patch_mode == "cfg":
|
||||
model.set_model_sampler_cfg_function(model_patch)
|
||||
elif patch_mode == "denoise_mask":
|
||||
model.set_model_denoise_mask_function(model_patch)
|
||||
elif patch_mode == "apply_model":
|
||||
model.set_model_unet_function_wrapper(model_patch)
|
||||
else:
|
||||
raise ValueError(f"Bad patch mode: {patch_mode}")
|
||||
return (model,)
|
||||
|
||||
|
||||
__all__ = (
|
||||
"SamplerNode",
|
||||
"GroupNode",
|
||||
|
||||
+84
-19
@@ -2,12 +2,67 @@ import gc
|
||||
import math
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import scipy
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .filtering import Filter, make_filter
|
||||
from .utils import scale_noise, fallback
|
||||
from .utils import fallback, scale_noise
|
||||
|
||||
try:
|
||||
from .triton_lsa import (
|
||||
assignments_to_indices,
|
||||
batch_linear_assignment,
|
||||
batch_linear_assignment_shuffled,
|
||||
)
|
||||
|
||||
HAVE_TRITON = True
|
||||
except Exception:
|
||||
HAVE_TRITON = False
|
||||
|
||||
|
||||
def linear_sum_assignment(
|
||||
cost: torch.Tensor,
|
||||
*,
|
||||
maximize: bool = False,
|
||||
use_triton: bool = False,
|
||||
split_batch: int = 0,
|
||||
**kwargs: dict,
|
||||
) -> tuple[np.ndarray, np.ndarray] | tuple[torch.Tensor, torch.Tensor]:
|
||||
if not use_triton or not HAVE_TRITON or not cost.is_cuda:
|
||||
cost = cost.half().cpu()
|
||||
return scipy.optimize.linear_sum_assignment(cost, maximize=maximize)
|
||||
ndim = cost.ndim
|
||||
orig_shape = cost.shape
|
||||
if ndim == 2:
|
||||
do_split = split_batch > 1 and all(
|
||||
(sz / split_batch).is_integer() for sz in orig_shape
|
||||
)
|
||||
if do_split:
|
||||
cost = cost.reshape(
|
||||
split_batch, orig_shape[0] // split_batch, orig_shape[1] // split_batch
|
||||
)
|
||||
else:
|
||||
cost = cost.unsqueeze(0)
|
||||
tqdm.write(
|
||||
f"TRITON LAP: maximize={maximize}, orig cost shape={orig_shape}, cost shape={cost.shape}, cost dtype={cost.dtype}",
|
||||
)
|
||||
if not cost.is_contiguous():
|
||||
cost = cost.contiguous()
|
||||
fun = (
|
||||
batch_linear_assignment
|
||||
if "generator" not in kwargs
|
||||
else batch_linear_assignment_shuffled
|
||||
)
|
||||
assignments = fun(cost, maximize=maximize, **kwargs)
|
||||
row_ind, col_ind = assignments_to_indices(assignments)
|
||||
if ndim == 2:
|
||||
row_ind = row_ind.reshape(-1, row_ind.shape[-1])
|
||||
col_ind = col_ind.reshape(-1, col_ind.shape[-1])
|
||||
# row_ind, col_ind = row_ind.squeeze(0), col_ind.squeeze(0)
|
||||
tqdm.write(f"Ran LAP kernel: {assignments.shape}, {row_ind.shape}, {col_ind.shape}")
|
||||
return row_ind, col_ind
|
||||
|
||||
|
||||
class ImmiscibleNoise(Filter):
|
||||
@@ -19,6 +74,12 @@ class ImmiscibleNoise(Filter):
|
||||
"maximize": False,
|
||||
"distance_scale": 0.0,
|
||||
"distance_scale_ref": None,
|
||||
"abs_mode": False,
|
||||
"abs_distance_mode": False,
|
||||
"use_triton": False,
|
||||
# Only honored in Triton mode.
|
||||
"split_batch": 0,
|
||||
"generator": None,
|
||||
}
|
||||
|
||||
def __call__(self, noise_sampler, x_ref, *, refs=None):
|
||||
@@ -83,7 +144,11 @@ class ImmiscibleNoise(Filter):
|
||||
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
|
||||
# Minimize latent-noise pairs over a batch
|
||||
batch = latent.shape[0]
|
||||
out_latent = fallback(out_latent, latent)
|
||||
ref_latent = ref_latent.detach().clone()
|
||||
if self.abs_mode:
|
||||
ref_latent = ref_latent.abs()
|
||||
latent = latent.abs()
|
||||
if self.distance_scale == 0:
|
||||
ref_latent_expanded = ref_latent.unsqueeze(1).expand(
|
||||
-1, batch, *ref_latent.shape[1:]
|
||||
@@ -92,30 +157,30 @@ class ImmiscibleNoise(Filter):
|
||||
ref_latent.shape[0], *latent.shape
|
||||
)
|
||||
dist = (ref_latent_expanded - latent_expanded) ** 2
|
||||
if self.abs_distance_mode:
|
||||
dist = dist.abs_()
|
||||
del ref_latent_expanded, latent_expanded
|
||||
dist = dist.mean(tuple(range(2, dist.dim())))
|
||||
dist = dist.mean(tuple(range(2, dist.ndim)))
|
||||
else:
|
||||
dist = torch.linalg.vector_norm(
|
||||
fallback(self.distance_scale_ref, self.distance_scale)
|
||||
* ref_latent.flatten(start_dim=1).unsqueeze(1)
|
||||
- self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0),
|
||||
dim=2,
|
||||
)
|
||||
dist = dist.half()
|
||||
distance_scale_ref = fallback(self.distance_scale_ref, self.distance_scale)
|
||||
dist = distance_scale_ref * ref_latent.flatten(start_dim=1).unsqueeze(
|
||||
1
|
||||
) - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0)
|
||||
if self.abs_distance_mode:
|
||||
dist = dist.abs_()
|
||||
dist = torch.linalg.vector_norm(dist, dim=2)
|
||||
try:
|
||||
assign_mat = scipy.optimize.linear_sum_assignment(
|
||||
dist.cpu(), maximize=self.maximize
|
||||
assign_mat = linear_sum_assignment(
|
||||
dist,
|
||||
maximize=self.maximize,
|
||||
use_triton=self.use_triton,
|
||||
split_batch=self.split_batch,
|
||||
generator=self.generator,
|
||||
)
|
||||
except ValueError as exc:
|
||||
tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}")
|
||||
return (
|
||||
None
|
||||
if return_idxs
|
||||
else fallback(out_latent, latent)[: ref_latent.shape[0]]
|
||||
)
|
||||
return (
|
||||
assign_mat if return_idxs else fallback(out_latent, latent)[assign_mat[1]]
|
||||
)
|
||||
return None if return_idxs else out_latent[: ref_latent.shape[0]]
|
||||
return assign_mat if return_idxs else out_latent[assign_mat[1]]
|
||||
|
||||
def immiscible_simple(
|
||||
self,
|
||||
|
||||
+112
-10
@@ -1,8 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
# from tqdm import tqdm
|
||||
|
||||
|
||||
class RestartScaleFactors(NamedTuple):
|
||||
latent_scale: float
|
||||
noise_scale: float
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
sigma_from: float | torch.Tensor,
|
||||
sigma_to: float | torch.Tensor,
|
||||
*,
|
||||
is_flow: bool,
|
||||
) -> RestartScaleFactors:
|
||||
if isinstance(sigma_from, torch.Tensor):
|
||||
sigma_from = sigma_from.max().item()
|
||||
if isinstance(sigma_to, torch.Tensor):
|
||||
sigma_to = sigma_to.max().item()
|
||||
if not is_flow:
|
||||
return cls(
|
||||
1.0,
|
||||
max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5,
|
||||
)
|
||||
alpha_from = 1.0 - sigma_from
|
||||
alpha_to = 1.0 - sigma_to
|
||||
if alpha_to <= 0:
|
||||
latent_scale = 0.0
|
||||
noise_scale = sigma_to
|
||||
else:
|
||||
latent_scale = alpha_to / alpha_from
|
||||
noise_scale = (
|
||||
max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5
|
||||
)
|
||||
return cls(latent_scale, noise_scale)
|
||||
|
||||
|
||||
class Restart:
|
||||
def __init__(self, *, s_noise=1.0, custom_noise=None, immiscible=False):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
s_noise=1.0,
|
||||
custom_noise=None,
|
||||
immiscible=False,
|
||||
normalized=True,
|
||||
normalize_dims: tuple[int, ...] | None = None,
|
||||
is_flow=False,
|
||||
):
|
||||
from .noise import ImmiscibleNoise
|
||||
|
||||
self.s_noise = s_noise
|
||||
@@ -10,6 +59,9 @@ class Restart:
|
||||
immiscible = ImmiscibleNoise(**immiscible)
|
||||
self.immiscible = immiscible
|
||||
self.custom_noise = custom_noise
|
||||
self.normalized = normalized
|
||||
self.normalize_dims = normalize_dims
|
||||
self.is_flow = is_flow
|
||||
|
||||
def get_noise_sampler(self, nsc):
|
||||
return nsc.make_caching_noise_sampler(
|
||||
@@ -30,23 +82,64 @@ class Restart:
|
||||
last_sigma = sigma
|
||||
return sigmas
|
||||
|
||||
def split_sigmas(self, sigmas):
|
||||
def split_sigmas(self, sigmas: torch.Tensor):
|
||||
prev_seg = None
|
||||
while len(sigmas) > 1:
|
||||
seg = self.get_segment(sigmas)
|
||||
sigmas = sigmas[len(seg) :]
|
||||
if prev_seg is not None and seg[0] > prev_seg[-1]:
|
||||
noise_scale = self.get_noise_scale(prev_seg[-1], seg[0])
|
||||
scale_factors = RestartScaleFactors.build(
|
||||
sigma_from=prev_seg[-1], sigma_to=seg[0], is_flow=self.is_flow
|
||||
)
|
||||
else:
|
||||
noise_scale = 0.0
|
||||
scale_factors = None
|
||||
prev_seg = seg
|
||||
yield (noise_scale, seg)
|
||||
yield (scale_factors, seg)
|
||||
|
||||
def get_noise_scale(self, s_min, s_max):
|
||||
def get_noise_scale(
|
||||
self, s_min: float | torch.Tensor, s_max: float | torch.Tensor
|
||||
) -> float:
|
||||
result = (s_max**2 - s_min**2) ** 0.5
|
||||
if isinstance(result, torch.Tensor):
|
||||
result = result.item()
|
||||
return result * self.s_noise
|
||||
return result.item()
|
||||
return result
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma_from: float,
|
||||
sigma_to: float,
|
||||
*,
|
||||
nsc,
|
||||
refs,
|
||||
scale_factors: RestartScaleFactors | None = None,
|
||||
in_place: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if self.is_flow:
|
||||
sigma_from = min(1.0, max(0.0, sigma_from))
|
||||
sigma_to = min(1.0, max(0.0, sigma_to))
|
||||
if sigma_from >= sigma_to:
|
||||
raise ValueError(
|
||||
f"sigma_from ({sigma_from:.4f}) must be less than sigma_to ({sigma_to:.4f})"
|
||||
)
|
||||
scale_factors = scale_factors or RestartScaleFactors.build(
|
||||
sigma_from, sigma_to, is_flow=self.is_flow
|
||||
)
|
||||
ns = self.get_noise_sampler(nsc)
|
||||
sigma_empty = nsc.min_sigma * 0
|
||||
noise = nsc.scale_noise(
|
||||
ns(sigma_empty + sigma_from, sigma_empty + sigma_to, refs=refs),
|
||||
normalized=self.normalized,
|
||||
normalize_dims=self.normalize_dims,
|
||||
)
|
||||
noise *= scale_factors.noise_scale * self.s_noise
|
||||
if scale_factors.latent_scale != 1.0:
|
||||
x = (
|
||||
x.mul_(scale_factors.latent_scale)
|
||||
if in_place
|
||||
else scale_factors.latent_scale * x
|
||||
)
|
||||
return noise.add_(x)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
|
||||
@@ -77,18 +170,27 @@ class Restart:
|
||||
raise ValueError("Schedule jump index out of range")
|
||||
sched_idx = item
|
||||
continue
|
||||
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
|
||||
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
|
||||
if sig_idx >= siglen or sig_idx < 0:
|
||||
break
|
||||
interval, jump = item
|
||||
chunk = siglist[sig_idx : sig_idx + interval + 1]
|
||||
if sched_frac != 0:
|
||||
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
|
||||
# print(f"{out} + {chunk}")
|
||||
out += chunk
|
||||
sig_idx += interval + jump
|
||||
if jump >= 0:
|
||||
sig_idx += 1
|
||||
sig_idx += interval + jump
|
||||
sched_idx += 1
|
||||
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
|
||||
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
|
||||
if sig_idx < siglen and sig_idx >= 0:
|
||||
out += siglist[sig_idx:]
|
||||
chunk = siglist[sig_idx:]
|
||||
if sched_frac != 0:
|
||||
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
|
||||
out += chunk
|
||||
if out[-1] > siglist[-1]:
|
||||
out.append(siglist[-1])
|
||||
return torch.tensor(out).to(sigmas)
|
||||
|
||||
+31
-21
@@ -1,13 +1,12 @@
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
|
||||
|
||||
from .filtering import FILTER_HANDLERS, FilterRefs
|
||||
from .model import OCSModel
|
||||
from .noise import NoiseSamplerCache
|
||||
from .substep_sampling import SamplerState
|
||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||
from .restart import Restart
|
||||
from .substep_merging import MERGE_SUBSTEPS_CLASSES
|
||||
from .substep_sampling import SamplerState
|
||||
|
||||
|
||||
def find_merge_sampler(merge_samplers, ss) -> object | None:
|
||||
@@ -48,11 +47,6 @@ def composable_sampler(
|
||||
restart_custom_noise = copts.get("restart_custom_noise")
|
||||
if isinstance(restart_custom_noise, str):
|
||||
restart_custom_noise = copts.get(f"restart_custom_noise_{restart_custom_noise}")
|
||||
restart = Restart(
|
||||
s_noise=restart_params.get("s_noise", 1.0),
|
||||
custom_noise=restart_custom_noise,
|
||||
immiscible=restart_params.get("immiscible", False),
|
||||
)
|
||||
|
||||
ss = SamplerState(
|
||||
OCSModel(
|
||||
@@ -72,6 +66,16 @@ def composable_sampler(
|
||||
reta=copts.get("reta", 1.0),
|
||||
disable_status=disable,
|
||||
)
|
||||
|
||||
restart = Restart(
|
||||
s_noise=restart_params.get("s_noise", 1.0),
|
||||
custom_noise=restart_custom_noise,
|
||||
immiscible=restart_params.get("immiscible", False),
|
||||
normalized=restart_params.get("normalized", True),
|
||||
normalize_dims=restart_params.get("normalize_dims"),
|
||||
is_flow=ss.model.is_rectified_flow,
|
||||
)
|
||||
|
||||
groups = copts["_groups"]
|
||||
merge_samplers = tuple(
|
||||
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
|
||||
@@ -85,17 +89,17 @@ def composable_sampler(
|
||||
)
|
||||
ss.noise = nsc
|
||||
sigma_chunks = (
|
||||
tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((0.0, sigmas),)
|
||||
tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((None, sigmas),)
|
||||
)
|
||||
step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks)
|
||||
ss.total_steps = step_count
|
||||
step = 0
|
||||
with trange(step_count, disable=ss.disable_status) as pbar:
|
||||
for noise_scale, chunk_sigmas in sigma_chunks:
|
||||
if step != 0 and noise_scale != 0:
|
||||
prev_refs = FilterRefs({
|
||||
f"pre_restart_{k}": v for k, v in ss.refs.items()
|
||||
})
|
||||
for chunk_idx, (scale_factors, chunk_sigmas) in enumerate(sigma_chunks):
|
||||
if step != 0 and scale_factors is not None:
|
||||
prev_refs = FilterRefs(
|
||||
{f"pre_restart_{k}": v for k, v in ss.refs.items()}
|
||||
)
|
||||
ss.sigmas = chunk_sigmas
|
||||
ss.update(0, step=step, substep=0)
|
||||
if step != 0:
|
||||
@@ -104,14 +108,20 @@ def composable_sampler(
|
||||
ss.hist.reset()
|
||||
for ms in merge_samplers:
|
||||
ms.reset()
|
||||
nsc.min_sigma, nsc.max_sigma = chunk_sigmas[-1], chunk_sigmas[0]
|
||||
if step != 0 and noise_scale != 0:
|
||||
restart_ns = restart.get_noise_sampler(nsc)
|
||||
x += nsc.scale_noise(
|
||||
restart_ns(nsc.min_sigma, nsc.max_sigma, refs=prev_refs | ss.refs),
|
||||
noise_scale,
|
||||
nsc.min_sigma, nsc.max_sigma = (
|
||||
chunk_sigmas[-1].clone(),
|
||||
chunk_sigmas[0].clone(),
|
||||
)
|
||||
if step != 0 and scale_factors is not None:
|
||||
x = restart.add_noise(
|
||||
x,
|
||||
sigma_from=sigma_chunks[chunk_idx - 1][1][-1].item(),
|
||||
sigma_to=chunk_sigmas[0].item(),
|
||||
scale_factors=scale_factors,
|
||||
nsc=nsc,
|
||||
refs=prev_refs | ss.refs,
|
||||
in_place=True,
|
||||
)
|
||||
del restart_ns
|
||||
del prev_refs
|
||||
for idx in range(len(chunk_sigmas) - 1):
|
||||
if idx > 0:
|
||||
|
||||
@@ -1,20 +1,19 @@
|
||||
import inspect
|
||||
import math
|
||||
import typing
|
||||
|
||||
import inspect
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
import comfy
|
||||
|
||||
from .. import filtering
|
||||
from .. import expression as expr
|
||||
from .. import filtering
|
||||
from ..utils import fallback
|
||||
from .base import (
|
||||
StepSamplerContext,
|
||||
SingleStepSampler,
|
||||
StepSamplerContext,
|
||||
registry,
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
import pytorch_wavelets as ptwav
|
||||
|
||||
@@ -531,7 +530,10 @@ class WeoonStep(SingleStepSampler):
|
||||
)
|
||||
denoised_new = self.wavelet_inverse(coeffs_out)
|
||||
if denoised_new.shape != x.shape:
|
||||
denoised_new = denoised_new.reshape(*x.shape)
|
||||
bi_elements = math.prod(x.shape[1:])
|
||||
denoised_new = denoised_new.reshape(x.shape[0], -1)[
|
||||
:, :bi_elements
|
||||
].reshape(*x.shape)
|
||||
x = self.blend(denoised_new, x, ratio)
|
||||
yield from self.result(x, sigma_up)
|
||||
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
import torch
|
||||
|
||||
import comfy
|
||||
import torch
|
||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||
from tqdm import tqdm
|
||||
|
||||
from .. import filtering
|
||||
from .base import (
|
||||
SingleStepSampler,
|
||||
DPMPPStepMixin,
|
||||
HistorySingleStepSampler,
|
||||
ReversibleSingleStepSampler,
|
||||
SingleStepSampler,
|
||||
registry,
|
||||
)
|
||||
|
||||
@@ -524,6 +525,147 @@ class RESMultistepStep(HistorySingleStepSampler, DPMPPStepMixin):
|
||||
yield from self.result(result, sigma_up, sigma_down=sigma_down)
|
||||
|
||||
|
||||
# SEEDS-2 - Stochastic Explicit Exponential Derivative-free Solvers (VP Data Prediction) stage 2.
|
||||
# arXiv: https://arxiv.org/abs/2305.14267 (NeurIPS 2023)
|
||||
# Implementation referenced from ComfyUI.
|
||||
class Seeds2Step(SingleStepSampler, DPMPPStepMixin):
|
||||
name = "seeds_2"
|
||||
self_noise = 3
|
||||
model_calls = 1
|
||||
allow_alt_cfgpp = False
|
||||
uses_alt_noise = True
|
||||
|
||||
def __init__(self, *args, r=0.5, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.r = r
|
||||
s2_options = self.options.get("seeds_2", {})
|
||||
sigma_blend_mode = s2_options.get("sigma_blend_mode", "lerp").strip()
|
||||
self.sigma_blend_function = (
|
||||
filtering.BLENDING_MODES[sigma_blend_mode]
|
||||
if sigma_blend_mode != "lerp"
|
||||
else torch.lerp
|
||||
)
|
||||
denoised_blend_mode = s2_options.get("denoised_blend_mode", "lerp").strip()
|
||||
self.denoised_blend_function = (
|
||||
filtering.BLENDING_MODES[denoised_blend_mode]
|
||||
if denoised_blend_mode != "lerp"
|
||||
else torch.lerp
|
||||
)
|
||||
self.disable_stage2_eta = bool(s2_options.get("disable_stage2_eta", False))
|
||||
stage2_stage1_noise_blend_mode = s2_options.get(
|
||||
"stage2_stage1_noise_blend_mode", "lerp"
|
||||
).strip()
|
||||
self.stage2_stage1_noise_blend_function = (
|
||||
filtering.BLENDING_MODES[stage2_stage1_noise_blend_mode]
|
||||
if stage2_stage1_noise_blend_mode != "lerp"
|
||||
else torch.lerp
|
||||
)
|
||||
self.stage2_stage1_noise_ratio = s2_options.get(
|
||||
"stage2_stage1_noise_ratio", 1.0
|
||||
)
|
||||
self.stage1_s_noise = s2_options.get("stage1_s_noise", 1.0)
|
||||
self.stage2_s_noise = s2_options.get("stage2_s_noise", 1.0)
|
||||
self.stage2_sigma_scale = s2_options.get("stage2_sigma_scale", 1.0)
|
||||
|
||||
def step(self, x: torch.Tensor):
|
||||
ss = self.ss
|
||||
sigma = ss.sigma.to(dtype=torch.float64)
|
||||
sigma_next = ss.sigma_next.to(dtype=torch.float64)
|
||||
denoised = ss.denoised
|
||||
|
||||
t_one = ss.sigma * 0 + 1.0
|
||||
|
||||
r, eta = self.r, self.get_dyn_eta()
|
||||
fac = 1 / (2 * r)
|
||||
lambda_s = ss.sigma_to_half_log_snr(sigma=sigma)
|
||||
lambda_t = ss.sigma_to_half_log_snr(sigma=sigma_next)
|
||||
h = lambda_t - lambda_s
|
||||
h_eta = h * (eta + 1.0)
|
||||
lambda_s_1 = self.sigma_blend_function(
|
||||
lambda_s.unsqueeze(0), lambda_t.unsqueeze(0), r
|
||||
).squeeze(0)
|
||||
sigma_s_1 = ss.half_log_snr_to_sigma(lambda_s_1)
|
||||
|
||||
alpha_s_1 = sigma_s_1 * lambda_s_1.exp()
|
||||
alpha_t = sigma_next * lambda_t.exp()
|
||||
|
||||
s1_x_mult = sigma_s_1 / sigma * (-r * h * eta).exp()
|
||||
s1_denoised_mult = alpha_s_1 * (-r * h_eta).expm1()
|
||||
x_2 = (
|
||||
s1_x_mult.to(dtype=x.dtype) * x
|
||||
- s1_denoised_mult.to(dtype=x.dtype) * denoised
|
||||
)
|
||||
if eta != 0:
|
||||
s1_noise_mult = (-2 * r * h * eta).expm1().neg().sqrt()
|
||||
sde_noise1 = yield from self.result(
|
||||
x_2 * 0,
|
||||
s1_noise_mult.to(dtype=x.dtype),
|
||||
sigma=ss.sigma,
|
||||
sigma_next=sigma_s_1.to(dtype=x.dtype),
|
||||
noise_sampler=self.alt_noise_sampler,
|
||||
final=False,
|
||||
)
|
||||
x_2 += sde_noise1 * (sigma_s_1.to(dtype=x.dtype) * self.stage1_s_noise)
|
||||
|
||||
denoised_2 = self.call_model(
|
||||
x_2, (sigma_s_1 * self.stage2_sigma_scale).to(dtype=x.dtype), call_index=1
|
||||
).denoised
|
||||
denoised_d = self.denoised_blend_function(denoised, denoised_2, fac)
|
||||
|
||||
if self.disable_stage2_eta:
|
||||
eta = 0.0
|
||||
h_eta = h
|
||||
|
||||
s2_x_mult = sigma_next / sigma * (-h * eta).exp()
|
||||
s2_denoised_mult = alpha_t * h_eta.neg().expm1()
|
||||
x_curr = s2_x_mult.to(dtype=x.dtype) * x
|
||||
x_curr -= s2_denoised_mult.to(dtype=x.dtype) * denoised_d
|
||||
|
||||
if eta == 0:
|
||||
return (yield from self.result(x_curr))
|
||||
|
||||
s2_s1_nr = self.stage2_stage1_noise_ratio
|
||||
|
||||
segment_factor = ((r - 1.0) * h * eta).to(dtype=x.dtype)
|
||||
s2_noise_mult = (segment_factor * 2.0).expm1().neg() ** 0.5
|
||||
sde_noise2_raw = yield from self.result(
|
||||
x_curr * 0,
|
||||
t_one,
|
||||
sigma=sigma_s_1.to(dtype=x.dtype),
|
||||
sigma_next=ss.sigma_next,
|
||||
final=False,
|
||||
)
|
||||
sde_noise2 = sde_noise2_raw * s2_noise_mult.to(dtype=x.dtype)
|
||||
|
||||
if s2_s1_nr != 1.0:
|
||||
# print(
|
||||
# f"\n\nBLENDING: {s2_s1_nr:.4f}, {s1_noise_mult.item():.4f}, {s2_noise_mult.item():.4f}"
|
||||
# )
|
||||
sde_noise1 = self.stage2_stage1_noise_blend_function(
|
||||
(
|
||||
yield from self.result(
|
||||
x_curr * 0,
|
||||
t_one,
|
||||
sigma=sigma_s_1.to(dtype=x.dtype),
|
||||
sigma_next=ss.sigma_next,
|
||||
final=False,
|
||||
)
|
||||
)
|
||||
* s1_noise_mult.to(dtype=x.dtype),
|
||||
sde_noise1,
|
||||
s2_s1_nr,
|
||||
)
|
||||
|
||||
sde_noise1 *= segment_factor.exp()
|
||||
sde_noise2 += sde_noise1
|
||||
sde_noise2 *= ss.sigma_next * self.stage2_s_noise
|
||||
x_curr += sde_noise2
|
||||
|
||||
yield from self.result(
|
||||
x_curr, noise_scale=ss.sigma * 0, sigma_down=ss.sigma_next
|
||||
)
|
||||
|
||||
|
||||
registry.add(
|
||||
DEISStep,
|
||||
DPMPP2MSDEStep,
|
||||
@@ -538,4 +680,5 @@ registry.add(
|
||||
DPM2Step,
|
||||
DPMPP2SStep,
|
||||
RESMultistepStep,
|
||||
Seeds2Step,
|
||||
)
|
||||
|
||||
@@ -10,7 +10,7 @@ class PingPongStep(SingleStepSampler):
|
||||
super().__init__(*args, **kwargs)
|
||||
pingpong_options = self.options.pop("pingpong", {})
|
||||
self.pingpong_start_step = pingpong_options.get("start_step", 0)
|
||||
self.pingpong_end_step = pingpong_options.get("end_step", 0)
|
||||
self.pingpong_end_step = pingpong_options.get("end_step", 9999)
|
||||
|
||||
def step(self, x):
|
||||
ss = self.ss
|
||||
|
||||
+16
-11
@@ -5,8 +5,7 @@ import tqdm
|
||||
|
||||
from . import expression as expr
|
||||
from . import utils
|
||||
|
||||
from .filtering import make_filter, FilterRefs, FILTER_HANDLERS
|
||||
from .filtering import FILTER_HANDLERS, FilterRefs, make_filter
|
||||
from .noise import ImmiscibleNoise
|
||||
from .restart import Restart
|
||||
from .step_samplers import STEP_SAMPLERS
|
||||
@@ -451,6 +450,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
s_noise=restart.get("s_noise", 1.0),
|
||||
custom_noise=restart_custom_noise,
|
||||
immiscible=restart.get("immiscible", False),
|
||||
is_flow=ss.model.is_rectified_flow,
|
||||
)
|
||||
|
||||
def make_schedule(self, ss):
|
||||
@@ -505,10 +505,13 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
if subss.idx >= max_idx:
|
||||
break
|
||||
if last_down is not None and last_down < ss.sigma_next:
|
||||
restart_ns = self.restart.get_noise_sampler(ss.noise)
|
||||
x += ss.noise.scale_noise(
|
||||
restart_ns(last_down, ss.sigma_next, refs=ss.refs),
|
||||
self.restart.get_noise_scale(last_down, ss.sigma_next),
|
||||
x = self.restart.add_noise(
|
||||
x,
|
||||
sigma_from=last_down.item(),
|
||||
sigma_to=ss.sigma_next.item(),
|
||||
nsc=nsc,
|
||||
refs=ss.refs,
|
||||
in_place=True,
|
||||
)
|
||||
pbar.update(0)
|
||||
return x
|
||||
@@ -653,11 +656,13 @@ class PingpongMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
sigma_next,
|
||||
immiscible=fallback(self.immiscible, ss.noise.immiscible),
|
||||
)
|
||||
noise_refs = ss.refs | FilterRefs({
|
||||
"orig_x": orig_x,
|
||||
"x": x,
|
||||
"denoised": synth_denoised,
|
||||
})
|
||||
noise_refs = ss.refs | FilterRefs(
|
||||
{
|
||||
"orig_x": orig_x,
|
||||
"x": x,
|
||||
"denoised": synth_denoised,
|
||||
}
|
||||
)
|
||||
noise = (
|
||||
noise_sampler(sigma, sigma_next, refs=noise_refs) * self.pingpong_s_noise
|
||||
)
|
||||
|
||||
+92
-8
@@ -1,5 +1,6 @@
|
||||
import torch
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||
|
||||
from .filtering import FilterRefs
|
||||
@@ -7,6 +8,13 @@ from .model import History
|
||||
from .utils import fallback
|
||||
|
||||
|
||||
class AncestralRatios(NamedTuple):
|
||||
alpha_t: torch.Tensor
|
||||
alpha_s: torch.Tensor
|
||||
sigma_up: torch.Tensor
|
||||
sigma_down: torch.Tensor
|
||||
|
||||
|
||||
class Items:
|
||||
def __init__(self, items=None):
|
||||
self.items = [] if items is None else items
|
||||
@@ -141,6 +149,10 @@ class SamplerState:
|
||||
self.substep = 0
|
||||
self.total_steps = len(sigmas) - 1
|
||||
self.cfg_scale_override = cfg_scale_override
|
||||
self.is_flow = self.model.is_rectified_flow
|
||||
self.offset_sigma = (
|
||||
model.model_sampling.percent_to_sigma(1e-04) if self.is_flow else None
|
||||
)
|
||||
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
|
||||
|
||||
@property
|
||||
@@ -171,6 +183,23 @@ class SamplerState:
|
||||
def d(self):
|
||||
return self.hcur.d
|
||||
|
||||
# These two functions referenced from ComfyUI.
|
||||
def sigma_to_half_log_snr(
|
||||
self, *, sigma: torch.Tensor | None = None, idx: int | None = None
|
||||
) -> torch.Tensor:
|
||||
if sigma is None and idx is None:
|
||||
sigma = self.sigma
|
||||
else:
|
||||
sigma = sigma if sigma is not None else self.sigmas[idx]
|
||||
if not self.is_flow:
|
||||
return sigma.log().neg_()
|
||||
if sigma.max() >= 1.0:
|
||||
sigma = sigma * 0.0 + self.offset_sigma
|
||||
return sigma.logit().neg_()
|
||||
|
||||
def half_log_snr_to_sigma(self, half_log_snr: torch.Tensor) -> torch.Tensor:
|
||||
return (torch.sigmoid if self.is_flow else torch.exp)(half_log_snr.neg())
|
||||
|
||||
def update(self, idx=None, step=None, substep=None):
|
||||
idx = self.idx if idx is None else idx
|
||||
self.idx = idx
|
||||
@@ -185,6 +214,59 @@ class SamplerState:
|
||||
self.substep = substep
|
||||
self.refs = FilterRefs.from_ss(self)
|
||||
|
||||
def get_ancestral_step_ext(
|
||||
self,
|
||||
*,
|
||||
sigma: torch.Tensor | None = None,
|
||||
sigma_next: torch.Tensor | None = None,
|
||||
eta: float = 1.0,
|
||||
retry_increment: int = 0,
|
||||
):
|
||||
sigma = fallback(sigma, self.sigma)
|
||||
sigma_next = fallback(sigma_next, self.sigma_next)
|
||||
sigma_empty = sigma_next * 0.0
|
||||
|
||||
def get_noeta_ratios():
|
||||
return AncestralRatios(
|
||||
alpha_t=sigma_empty + 1.0,
|
||||
alpha_s=sigma_empty + 1.0,
|
||||
sigma_up=sigma_empty.clone(),
|
||||
sigma_down=sigma_next.clone(),
|
||||
)
|
||||
|
||||
if eta <= 0 or sigma_next.max().item() <= 1e-08:
|
||||
return get_noeta_ratios()
|
||||
orig_dtype = sigma.dtype
|
||||
sigma = sigma.to(dtype=torch.float64)
|
||||
sigma_next = sigma_next.to(dtype=torch.float64)
|
||||
alpha_s = sigma * self.sigma_to_half_log_snr(sigma=sigma).exp()
|
||||
alpha_t = sigma_next * self.sigma_to_half_log_snr(sigma=sigma_next).exp()
|
||||
adj_sigma = sigma / alpha_s
|
||||
adj_sigma_next = sigma_next / alpha_t
|
||||
sd = su = None
|
||||
while eta > 0:
|
||||
sd, su = (
|
||||
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
|
||||
for v in get_ancestral_step(adj_sigma, adj_sigma_next, eta=eta)
|
||||
)
|
||||
if sd > 0 and su > 0:
|
||||
break
|
||||
else:
|
||||
sd = su = None
|
||||
if retry_increment <= 0:
|
||||
break
|
||||
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
|
||||
eta -= retry_increment
|
||||
if sd is None or su is None:
|
||||
return get_noeta_ratios()
|
||||
sd = alpha_t * sd
|
||||
return AncestralRatios(
|
||||
alpha_t=alpha_t.to(dtype=orig_dtype),
|
||||
alpha_s=alpha_s.to(dtype=orig_dtype),
|
||||
sigma_up=su.to(dtype=orig_dtype),
|
||||
sigma_down=sd.to(dtype=orig_dtype),
|
||||
)
|
||||
|
||||
def get_ancestral_step(
|
||||
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
|
||||
):
|
||||
@@ -265,13 +347,15 @@ class SamplerState:
|
||||
preview = (hi.x - hi.denoised) * 0.1 + hi.denoised
|
||||
else:
|
||||
preview = hi.denoised
|
||||
return self.callback_({
|
||||
"x": hi.x,
|
||||
"i": self.step,
|
||||
"sigma": hi.sigma,
|
||||
"sigma_hat": hi.sigma,
|
||||
"denoised": preview,
|
||||
})
|
||||
return self.callback_(
|
||||
{
|
||||
"x": hi.x,
|
||||
"i": self.step,
|
||||
"sigma": hi.sigma,
|
||||
"sigma_hat": hi.sigma,
|
||||
"denoised": preview,
|
||||
}
|
||||
)
|
||||
|
||||
def reset(self):
|
||||
self.hist.reset()
|
||||
|
||||
@@ -0,0 +1,386 @@
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({"num_warps": 4, "num_stages": 2}, num_warps=4, num_stages=2),
|
||||
triton.Config({"num_warps": 8, "num_stages": 2}, num_warps=8, num_stages=2),
|
||||
triton.Config({"num_warps": 4, "num_stages": 3}, num_warps=4, num_stages=3),
|
||||
triton.Config({"num_warps": 8, "num_stages": 3}, num_warps=8, num_stages=3),
|
||||
],
|
||||
key=[
|
||||
"B",
|
||||
"R",
|
||||
"C",
|
||||
"BLOCK_SIZE",
|
||||
], # Retune if matrix dimensions change significantly
|
||||
)
|
||||
@triton.jit
|
||||
def auction_lap_kernel(
|
||||
cost_ptr,
|
||||
assign_ptr,
|
||||
stride_b,
|
||||
stride_r,
|
||||
stride_c,
|
||||
stride_assign_b,
|
||||
stride_assign_r,
|
||||
B: tl.constexpr,
|
||||
R: tl.constexpr,
|
||||
C: tl.constexpr,
|
||||
epsilon,
|
||||
max_iter,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
|
||||
cost_base = cost_ptr + pid * stride_b
|
||||
assign_base = assign_ptr + pid * stride_assign_b
|
||||
|
||||
offs = tl.arange(0, BLOCK_SIZE)
|
||||
col_mask = offs < C
|
||||
|
||||
# Prices and Owners in SRAM/Registers
|
||||
prices = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
|
||||
owners = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
|
||||
row_to_col = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
|
||||
|
||||
iter_idx = 0
|
||||
unassigned_count = R
|
||||
|
||||
loop_continue = tl.full([], 1, dtype=tl.int1)
|
||||
|
||||
# Loop condition:
|
||||
# 1. unassigned_count > 0: Logic handled inside, but we need a break mechanism
|
||||
# 2. iter_idx < max_iter: Safety break
|
||||
# 3. loop_continue: Did we make progress last time?
|
||||
|
||||
while unassigned_count > 0 and iter_idx < max_iter and loop_continue:
|
||||
# Reset progress flag
|
||||
# loop_continue &= False
|
||||
loop_continue = tl.full([], 0, dtype=tl.int1)
|
||||
|
||||
# Gauss-Seidel pass over all rows
|
||||
for i in tl.range(0, R):
|
||||
# Check if row i is unassigned
|
||||
curr_c = tl.sum(tl.where(offs == i, row_to_col, 0))
|
||||
|
||||
if curr_c == -1:
|
||||
# Load costs
|
||||
row_cost_ptr = cost_base + i * stride_r + offs
|
||||
row_costs = tl.load(row_cost_ptr, mask=col_mask, other=-torch.inf)
|
||||
|
||||
# Net Value
|
||||
values = row_costs - prices
|
||||
|
||||
# Find Best
|
||||
best_val, best_idx = tl.max(values, axis=0, return_indices=True)
|
||||
|
||||
# CRITICAL: Only proceed if this is a valid edge (not -inf)
|
||||
if best_val > -torch.inf:
|
||||
# We have a valid move, so we continue the outer loop
|
||||
loop_continue = tl.full([], 1, dtype=tl.int1)
|
||||
|
||||
# Find Second Best
|
||||
mask_not_best = (offs != best_idx) & col_mask
|
||||
vals_no_best = tl.where(mask_not_best, values, -torch.inf)
|
||||
second_best_val = tl.max(vals_no_best, axis=0)
|
||||
|
||||
# Compute Bid
|
||||
bid = best_val - second_best_val + epsilon
|
||||
|
||||
# Update Price
|
||||
prices = tl.where(offs == best_idx, prices + bid, prices)
|
||||
|
||||
# Update Owners
|
||||
prev_owner = tl.sum(tl.where(offs == best_idx, owners, 0))
|
||||
|
||||
if prev_owner != -1:
|
||||
# Kick out previous owner
|
||||
row_to_col = tl.where(offs == prev_owner, -1, row_to_col)
|
||||
unassigned_count += 1
|
||||
|
||||
# Assign to current row
|
||||
owners = tl.where(offs == best_idx, i, owners)
|
||||
row_to_col = tl.where(offs == i, best_idx, row_to_col)
|
||||
unassigned_count -= 1
|
||||
|
||||
iter_idx += 1
|
||||
|
||||
# Store Result
|
||||
store_offs = tl.arange(0, BLOCK_SIZE)
|
||||
store_mask = store_offs < R
|
||||
tl.store(assign_base + store_offs, row_to_col, mask=store_mask)
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Python Helpers
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def rescale_simple(
|
||||
t: torch.Tensor,
|
||||
target_min: float = 0.0,
|
||||
target_max: float = 1.0,
|
||||
*,
|
||||
start_dim: int = 1,
|
||||
eps: float = 1e-07,
|
||||
) -> torch.Tensor:
|
||||
width = target_max - target_min
|
||||
if width == 0.0:
|
||||
return torch.zeros_like(t)
|
||||
orig_shape = t.shape
|
||||
t = t.flatten(start_dim=start_dim)
|
||||
min_val, max_val = t.aminmax(dim=-1, keepdim=True)
|
||||
normalized = t - min_val
|
||||
normalized /= (max_val - min_val).add_(eps)
|
||||
normalized *= width
|
||||
if target_min != 0.0:
|
||||
normalized += target_min
|
||||
return normalized.clamp_(target_min, target_max).reshape(orig_shape)
|
||||
|
||||
|
||||
def _greedy_fill_missing(assignments: torch.Tensor, C: int) -> None:
|
||||
"""
|
||||
Fills unassigned rows (-1) in the assignments tensor with available columns.
|
||||
This acts as a fallback when the Auction algorithm hits max_iter without
|
||||
full convergence.
|
||||
|
||||
Args:
|
||||
assignments: Tensor of shape (B, R) containing col indices or -1.
|
||||
C: Total number of columns available.
|
||||
"""
|
||||
# Identify which batch items have unassigned rows
|
||||
# This is usually a very small subset (e.g., < 1% of the batch)
|
||||
problem_batches = (assignments == -1).any(dim=1).nonzero().flatten()
|
||||
|
||||
if problem_batches.numel() == 0:
|
||||
return
|
||||
|
||||
device = assignments.device
|
||||
|
||||
# Iterate only over the problematic batch items
|
||||
# (Looping is acceptable here as B_subset is typically tiny)
|
||||
for b_idx in problem_batches:
|
||||
# 1. Find which rows are missing an assignment
|
||||
row_mask = assignments[b_idx] == -1
|
||||
missing_rows = row_mask.nonzero().flatten()
|
||||
n_needed = missing_rows.shape[0]
|
||||
|
||||
# 2. Find which columns are already used
|
||||
used_cols = assignments[b_idx][~row_mask]
|
||||
|
||||
# 3. Find free columns (Set difference: All - Used)
|
||||
# Create a boolean mask of all columns, then mark used ones as False
|
||||
# efficient on GPU for mid-sized C
|
||||
col_mask = torch.ones(C, device=device, dtype=torch.bool)
|
||||
col_mask[used_cols.long()] = False
|
||||
|
||||
free_cols = col_mask.nonzero().flatten()
|
||||
|
||||
# 4. Assign the first N free columns to the N missing rows
|
||||
# Since R <= C in this context (due to transpose logic in wrapper),
|
||||
# free_cols.numel() is guaranteed to be >= n_needed.
|
||||
assignments[b_idx, missing_rows] = free_cols[:n_needed].to(assignments.dtype)
|
||||
|
||||
|
||||
def batch_linear_assignment(
|
||||
cost_matrix: torch.Tensor,
|
||||
*,
|
||||
maximize: bool = False,
|
||||
max_iter: int | None = None,
|
||||
fill_missing: bool = True,
|
||||
rescale_costs: tuple[float, float] | None = (0.0, 1.0),
|
||||
invert_costs_mode: bool = True,
|
||||
eps: float = 1e-3,
|
||||
):
|
||||
if cost_matrix.ndim != 3:
|
||||
raise ValueError("Cost matrix must be (B, R, C)")
|
||||
if not cost_matrix.is_cuda:
|
||||
raise ValueError("Cost matrix must be a CUDA tensor")
|
||||
if not cost_matrix.is_contiguous():
|
||||
raise ValueError("Cost matrix must be contiguous")
|
||||
|
||||
B, R, C = cost_matrix.shape
|
||||
device = cost_matrix.device
|
||||
|
||||
# 1. Handle Rectangular Matrices
|
||||
# The Auction algorithm assigns Rows -> Cols.
|
||||
# It naturally handles R <= C (finding best col for every row).
|
||||
# If R > C, we must transpose to match Cols -> Rows, then invert the result.
|
||||
if R > C:
|
||||
transposed = True
|
||||
cost_matrix = cost_matrix.mT.contiguous()
|
||||
# Swap R and C for the kernel execution
|
||||
R, C = C, R
|
||||
else:
|
||||
transposed = False
|
||||
|
||||
if rescale_costs is not None:
|
||||
cost_matrix = rescale_simple(cost_matrix, *rescale_costs)
|
||||
# Note: We use float32 for atomic compatibility and speed
|
||||
cost_matrix = cost_matrix.to(torch.float32, copy=rescale_costs is None)
|
||||
|
||||
if not maximize:
|
||||
# Maximize (Value - Price) -> Minimize Cost
|
||||
cost_matrix = cost_matrix.neg_()
|
||||
if invert_costs_mode and rescale_costs is not None:
|
||||
cost_matrix += sum(rescale_costs)
|
||||
|
||||
assignments = torch.full(
|
||||
(B, R),
|
||||
-1,
|
||||
device=device,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
|
||||
max_dim = max(R, C)
|
||||
BLOCK_SIZE = max(32, triton.next_power_of_2(max_dim))
|
||||
|
||||
# Safety limit
|
||||
max_iter = max_iter if max_iter is not None else int(max(2000, R * C))
|
||||
|
||||
grid = (B,)
|
||||
|
||||
auction_lap_kernel[grid](
|
||||
cost_matrix,
|
||||
assignments,
|
||||
cost_matrix.stride(0),
|
||||
cost_matrix.stride(1),
|
||||
cost_matrix.stride(2),
|
||||
assignments.stride(0),
|
||||
assignments.stride(1),
|
||||
B,
|
||||
R,
|
||||
C,
|
||||
eps,
|
||||
max_iter,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
if fill_missing:
|
||||
_greedy_fill_missing(assignments, C)
|
||||
|
||||
assignments = assignments.long()
|
||||
|
||||
if not transposed:
|
||||
return assignments
|
||||
|
||||
# 2. Post-process Rectangular Results
|
||||
# We computed Col -> Row. We need Row -> Col.
|
||||
# assignments shape is currently (B, Original_Cols)
|
||||
# We want output shape (B, Original_Rows)
|
||||
|
||||
real_rows = C # C is the 'large' dimension (Original Rows)
|
||||
output = torch.full((B, real_rows), -1, device=device, dtype=torch.long)
|
||||
|
||||
# Create indices for the scatter source
|
||||
# We want: output[row_idx] = col_idx
|
||||
# Currently we have: assignments[col_idx] = row_idx
|
||||
src_col_indices = torch.arange(R, device=device).unsqueeze(0).expand(B, R)
|
||||
|
||||
# We use scatter. index=assignments (the rows), src=col_indices
|
||||
# To handle -1s in assignments, we clamp to 0 and then mask the result
|
||||
safe_assigns = assignments.clamp(min=0)
|
||||
output.scatter_(1, safe_assigns, src_col_indices)
|
||||
|
||||
# Cleanup: Any row that wasn't targeted by the scatter should be -1
|
||||
# The scatter might have written to index 0 if assignment was -1
|
||||
# Re-verify logic:
|
||||
for b in range(B):
|
||||
valid_mask = assignments[b] >= 0
|
||||
# Reset output
|
||||
output[b].fill_(-1)
|
||||
# Only write valid mappings
|
||||
# output[b, row_id] = col_id
|
||||
output[b, assignments[b, valid_mask]] = src_col_indices[b, valid_mask]
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def assignments_to_indices(
|
||||
assignments: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Converts a dense assignment tensor (from Triton/Hungraian) to
|
||||
batched SciPy-style indices.
|
||||
|
||||
Args:
|
||||
assignments (torch.Tensor): Shape (B, R). Values are col indices or -1.
|
||||
|
||||
Returns:
|
||||
row_ind (torch.Tensor): Shape (B, K) where K = min(R, C).
|
||||
col_ind (torch.Tensor): Shape (B, K).
|
||||
"""
|
||||
B, R = assignments.shape
|
||||
device = assignments.device
|
||||
|
||||
# 1. Create a mask of valid assignments (values >= 0)
|
||||
# In a rectangular assignment, the number of valid matches
|
||||
# is always min(Rows, Cols).
|
||||
mask = assignments >= 0
|
||||
|
||||
# 2. Extract Column Indices
|
||||
# We select the values from the assignment tensor that are valid.
|
||||
# We reshape to (B, -1) to preserve the batch dimension.
|
||||
col_ind = assignments[mask].view(B, -1)
|
||||
|
||||
# 3. Extract Row Indices
|
||||
# We need a grid of row indices [0, 1, 2, ... R-1] repeated B times
|
||||
row_grid = (
|
||||
torch.arange(R, device=device, dtype=assignments.dtype)
|
||||
.unsqueeze(0)
|
||||
.expand(B, R)
|
||||
)
|
||||
row_ind = row_grid[mask].view(B, -1)
|
||||
|
||||
return row_ind, col_ind
|
||||
|
||||
|
||||
def batch_linear_assignment_shuffled(
|
||||
cost_matrix: torch.Tensor,
|
||||
*args,
|
||||
**kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
generator = kwargs.pop("generator", None)
|
||||
# cost_matrix: [B, R, C]
|
||||
B, R = cost_matrix.shape[:2]
|
||||
|
||||
# 1. Generate a random permutation for the rows
|
||||
# We use one perm for the whole batch for efficiency,
|
||||
# or you can do it per-batch-item if B is small and quality is critical.
|
||||
# Here we shuffle all rows commonly.
|
||||
perm = torch.randperm(
|
||||
R,
|
||||
device=cost_matrix.device,
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
# 2. Shuffle the input (Row dimension is dim 1)
|
||||
# This creates a shuffled view/copy of the cost matrix
|
||||
shuffled_cost = cost_matrix[:, perm, :]
|
||||
|
||||
# 3. Run the Solver
|
||||
shuffled_assignments = batch_linear_assignment(
|
||||
shuffled_cost,
|
||||
*args,
|
||||
**kwargs,
|
||||
) # Returns [B, R]
|
||||
|
||||
# 4. Un-shuffle the results
|
||||
# We need to map the results back to their original row positions.
|
||||
# shuffled_assignments[b, i] corresponds to the row 'perm[i]'
|
||||
# We want final_assignments[b, perm[i]] = shuffled_assignments[b, i]
|
||||
|
||||
# Create the inverse permutation or just scatter back
|
||||
final_assignments = torch.empty_like(shuffled_assignments)
|
||||
|
||||
# Expand perm for the batch: [B, R]
|
||||
batch_perm = perm.unsqueeze(0).expand(B, R)
|
||||
|
||||
# Scatter the results back to original positions
|
||||
# dim=1, index=batch_perm, src=shuffled_assignments
|
||||
final_assignments.scatter_(1, batch_perm, shuffled_assignments)
|
||||
|
||||
return final_assignments
|
||||
+296
-55
@@ -1,7 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.k_diffusion.sampling import to_d
|
||||
|
||||
# def scale_noise_(
|
||||
@@ -39,25 +42,118 @@ from comfy.k_diffusion.sampling import to_d
|
||||
|
||||
|
||||
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, ...] = (-3, -2, -1),
|
||||
eps: float = 1e-08,
|
||||
) -> torch.Tensor:
|
||||
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)
|
||||
std = noise.std(dim=normalize_dims, keepdim=True)
|
||||
noise = noise / torch.where(std != 0.0, std, eps)
|
||||
noise -= noise.mean(dim=normalize_dims, keepdim=True)
|
||||
return noise if factor == 1.0 else noise.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 _quantile_norm_scaledown(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
*,
|
||||
dim,
|
||||
**_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)
|
||||
noiseabs = noise.abs()
|
||||
mv = noiseabs.max(dim=dim, keepdim=True).values.clamp(min=1e-06)
|
||||
return (
|
||||
noise
|
||||
if mv.sum().item() == 0
|
||||
else torch.where(noiseabs > nq, noise * (nq / mv), noise)
|
||||
)
|
||||
|
||||
|
||||
def _quantile_norm_wave(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
*,
|
||||
preserve_sign: bool = False,
|
||||
wave_function=torch.sin,
|
||||
pi_factor: float = 0.5,
|
||||
wrong_mode: bool = False,
|
||||
**_kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
if wrong_mode:
|
||||
multiplier = 1.0 / ((math.pi * pi_factor) / nq)
|
||||
else:
|
||||
multiplier = 1.0 / (nq / (math.pi * pi_factor))
|
||||
pos_mask = noise >= 0
|
||||
neg_mask = ~pos_mask
|
||||
result = torch.zeros_like(noise)
|
||||
result[pos_mask] = wave_function(noise.mul(multiplier))[pos_mask]
|
||||
result[neg_mask] = wave_function(noise.mul(multiplier))[neg_mask]
|
||||
result *= nq
|
||||
return result.copysign(noise) if preserve_sign else result
|
||||
|
||||
|
||||
def _quantile_norm_mode(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
*,
|
||||
dim: int | None,
|
||||
decimals=1,
|
||||
**_kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
return torch.where(
|
||||
noise.abs() > nq,
|
||||
noise.round(decimals=decimals).mode(dim=dim, keepdim=True).values,
|
||||
noise,
|
||||
)
|
||||
|
||||
|
||||
def _quantile_norm_replace(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
*,
|
||||
keep_sign: bool = False,
|
||||
avoid_sign: bool = False,
|
||||
count: int = 1,
|
||||
count_flipping: bool = False,
|
||||
**_kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
mask = noise.abs() <= nq
|
||||
candidates = noise[mask].flatten()
|
||||
n_candidates = candidates.numel()
|
||||
idxs = torch.arange(noise.numel()) % n_candidates
|
||||
cresult = candidates[idxs]
|
||||
if count < 2:
|
||||
candidates = cresult
|
||||
else:
|
||||
multiplier = 1.0 / count
|
||||
cresult = cresult * multiplier # noqa: PLR6104
|
||||
for i in range(1, count):
|
||||
cresult += (
|
||||
candidates[
|
||||
torch.roll(
|
||||
idxs,
|
||||
i if not count_flipping or (i % 2) == 0 else -i,
|
||||
dims=(-1,),
|
||||
)
|
||||
]
|
||||
* multiplier
|
||||
)
|
||||
candidates = cresult.reshape(noise.shape)
|
||||
if keep_sign or avoid_sign:
|
||||
candidates = candidates.copysign_(noise.neg() if avoid_sign else noise)
|
||||
return torch.where(mask, noise, candidates)
|
||||
|
||||
|
||||
quantile_handlers = {
|
||||
@@ -69,14 +165,66 @@ quantile_handlers = {
|
||||
noise.tanh().mul_(nq.abs()),
|
||||
noise,
|
||||
),
|
||||
"sigmoid": lambda noise, nq, **_kwargs: noise.sigmoid()
|
||||
.mul_(nq.abs())
|
||||
.copysign(noise),
|
||||
"sigmoid_keepsign": lambda noise, nq, **_kwargs: (
|
||||
noise.sigmoid().mul_(nq.abs()).copysign(noise)
|
||||
),
|
||||
"sigmoid": lambda noise, nq, **_kwargs: (
|
||||
noise.sigmoid().mul_(nq.abs() * 2).sub_(nq.abs())
|
||||
),
|
||||
"sigmoid_outliers": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
noise.sigmoid().mul_(nq.abs()).copysign(noise),
|
||||
noise,
|
||||
),
|
||||
"sin": partial(_quantile_norm_wave, wave_function=torch.sin),
|
||||
"sin_wholepi": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.sin,
|
||||
pi_factor=1.0,
|
||||
),
|
||||
"sin_keepsign": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.sin,
|
||||
preserve_sign=True,
|
||||
),
|
||||
"sin_wrong": partial(_quantile_norm_wave, wave_function=torch.sin, wrong_mode=True),
|
||||
"sin_wrong_wholepi": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.sin,
|
||||
pi_factor=1.0,
|
||||
wrong_mode=True,
|
||||
),
|
||||
"sin_wrong_keepsign": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.sin,
|
||||
preserve_sign=True,
|
||||
wrong_mode=True,
|
||||
),
|
||||
"cos": partial(_quantile_norm_wave, wave_function=torch.cos),
|
||||
"cos_wholepi": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.cos,
|
||||
pi_factor=1.0,
|
||||
),
|
||||
"cos_keepsign": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.cos,
|
||||
preserve_sign=True,
|
||||
),
|
||||
"cos_wrong": partial(_quantile_norm_wave, wave_function=torch.cos, wrong_mode=True),
|
||||
"cos_wrong_wholepi": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.cos,
|
||||
pi_factor=1.0,
|
||||
wrong_mode=True,
|
||||
),
|
||||
"cos_wrong_keepsign": partial(
|
||||
_quantile_norm_wave,
|
||||
wave_function=torch.cos,
|
||||
preserve_sign=True,
|
||||
wrong_mode=True,
|
||||
),
|
||||
"atan": lambda noise, nq, **_kwargs: noise.atan().mul_(nq.abs() / (math.pi / 2)),
|
||||
"tenth": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
noise * 0.1,
|
||||
@@ -93,6 +241,80 @@ quantile_handlers = {
|
||||
noise,
|
||||
0,
|
||||
),
|
||||
"mean": lambda noise, nq, *, dim, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
noise.mean(dim=dim, keepdim=True),
|
||||
noise,
|
||||
),
|
||||
"median": lambda noise, nq, *, dim, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
noise.median(dim=dim, keepdim=True).values,
|
||||
noise,
|
||||
),
|
||||
"mode_1dec": partial(_quantile_norm_mode, decimals=1),
|
||||
"mode_2dec": partial(_quantile_norm_mode, decimals=2),
|
||||
"replace": _quantile_norm_replace,
|
||||
"replace_keepsign": partial(_quantile_norm_replace, keep_sign=True),
|
||||
"replace_avoidsign": partial(_quantile_norm_replace, avoid_sign=True),
|
||||
"replace_2pt": partial(_quantile_norm_replace, count=2),
|
||||
"replace_3pt": partial(_quantile_norm_replace, count=3),
|
||||
"replace_2pt_flip": partial(_quantile_norm_replace, count=2, count_flipping=True),
|
||||
"replace_3pt_flip": partial(_quantile_norm_replace, count=3, count_flipping=True),
|
||||
"replace_2pt_keepsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=2,
|
||||
keep_sign=True,
|
||||
),
|
||||
"replace_3pt_keepsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=3,
|
||||
keep_sign=True,
|
||||
),
|
||||
"replace_2pt_flip_keepsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=2,
|
||||
count_flipping=True,
|
||||
keep_sign=True,
|
||||
),
|
||||
"replace_3pt_flip_keepsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=3,
|
||||
count_flipping=True,
|
||||
keep_sign=True,
|
||||
),
|
||||
"replace_2pt_avoidsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=2,
|
||||
avoid_sign=True,
|
||||
),
|
||||
"replace_3pt_avoidsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=3,
|
||||
avoid_sign=True,
|
||||
),
|
||||
"replace_2pt_flip_avoidsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=2,
|
||||
count_flipping=True,
|
||||
avoid_sign=True,
|
||||
),
|
||||
"replace_3pt_flip_avoidsign": partial(
|
||||
_quantile_norm_replace,
|
||||
count=3,
|
||||
count_flipping=True,
|
||||
avoid_sign=True,
|
||||
),
|
||||
"wrap": lambda noise, nq, **_kwargs: range_wrap(noise, -nq, nq),
|
||||
"wrap_keepsign": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
range_wrap(noise, -nq, nq).copysign_(noise),
|
||||
noise,
|
||||
),
|
||||
"wrap_avoidsign": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
range_wrap(noise, -nq, nq).copysign_(noise.neg()),
|
||||
noise,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@@ -100,48 +322,40 @@ quantile_handlers = {
|
||||
def quantile_normalize(
|
||||
noise: torch.Tensor,
|
||||
*,
|
||||
quantile: float = 0.75,
|
||||
quantile: float | tuple | list = 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,
|
||||
eps=1e-08,
|
||||
) -> torch.Tensor:
|
||||
if quantile is None or quantile <= 0 or quantile >= 1:
|
||||
if noise.numel() == 0:
|
||||
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",
|
||||
for q in quantile:
|
||||
noise = quantile_normalize(
|
||||
noise=noise,
|
||||
quantile=q,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
nq_fac=nq_fac,
|
||||
pow_fac=pow_fac,
|
||||
strategy=strategy,
|
||||
strategy_handler=strategy_handler,
|
||||
)
|
||||
return noise
|
||||
if quantile is None or quantile >= 1 or quantile <= -1:
|
||||
return noise
|
||||
centered = quantile < 0
|
||||
absquantile = abs(quantile)
|
||||
orig_shape = noise.shape
|
||||
if noise.ndim > 1 and flatten:
|
||||
flatnoise = noise.flatten(start_dim=dim)
|
||||
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)
|
||||
flatten = False
|
||||
flatnoise = noise
|
||||
handler = (
|
||||
quantile_handlers.get(strategy)
|
||||
if strategy_handler is None
|
||||
@@ -149,18 +363,45 @@ def quantile_normalize(
|
||||
)
|
||||
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()
|
||||
if not centered:
|
||||
nq = torch.quantile(
|
||||
flatnoise.abs(),
|
||||
quantile,
|
||||
dim=-1 if flatten else dim,
|
||||
keepdim=True,
|
||||
)
|
||||
return noise
|
||||
nq = nq.mul_(nq_fac).add_(eps)
|
||||
# print(f"\nNQ: {nq}")
|
||||
noise = handler(
|
||||
flatnoise,
|
||||
nq,
|
||||
orig_noise=noise,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
)
|
||||
else:
|
||||
absnoise = flatnoise.abs()
|
||||
maxabs = absnoise.amax(dim=-1 if flatten else dim, keepdim=True)
|
||||
proxy = flatnoise.sign().mul_(maxabs - absnoise)
|
||||
nq_proxy = torch.quantile(
|
||||
proxy.abs(),
|
||||
absquantile,
|
||||
dim=-1 if flatten else dim,
|
||||
keepdim=True,
|
||||
)
|
||||
nq_proxy = nq_proxy.mul_(nq_fac).add_(eps)
|
||||
# print(f"\nNQ proxy: {nq_proxy}")
|
||||
out_proxy = handler(
|
||||
proxy,
|
||||
nq_proxy,
|
||||
orig_noise=noise,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
)
|
||||
noise = out_proxy.sign().mul_(maxabs - out_proxy.abs())
|
||||
if pow_fac not in {0.0, 1.0}:
|
||||
noise = noise.abs().pow_(pow_fac).copysign(noise)
|
||||
return noise if noise.shape == orig_shape else noise.reshape(orig_shape)
|
||||
|
||||
|
||||
# def scale_noise(
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from .utils import fallback
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
try:
|
||||
import pytorch_wavelets as ptwav
|
||||
import pywt
|
||||
|
||||
HAVE_WAVELETS = True
|
||||
except ImportError:
|
||||
ptwav = None
|
||||
pywt = None
|
||||
HAVE_WAVELETS = False
|
||||
|
||||
|
||||
class Wavelet:
|
||||
DEFAULT_MODE = "symmetric"
|
||||
DEFAULT_LEVEL = 3
|
||||
DEFAULT_WAVE = "db4"
|
||||
DEFAULT_USE_1D_DWT = False
|
||||
DEFAULT_USE_DTCWT = False
|
||||
DEFAULT_QSHIFT = "qshift_a"
|
||||
DEFAULT_BIORT = "near_sym_a"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
wave: str = DEFAULT_WAVE,
|
||||
level: int = DEFAULT_LEVEL,
|
||||
mode: str = DEFAULT_MODE,
|
||||
use_1d_dwt: bool = DEFAULT_USE_1D_DWT,
|
||||
use_dtcwt: bool = DEFAULT_USE_DTCWT,
|
||||
biort: str = DEFAULT_BIORT,
|
||||
qshift: str = DEFAULT_QSHIFT,
|
||||
inv_wave: str | None = None,
|
||||
inv_mode: str | None = None,
|
||||
inv_biort: str | None = None,
|
||||
inv_qshift=None,
|
||||
device: str | torch.device | None = None,
|
||||
):
|
||||
if not HAVE_WAVELETS:
|
||||
raise RuntimeError(
|
||||
"Wavelet noise requires the pytorch_wavelets package to be installed in your Python environment",
|
||||
)
|
||||
inv_wave = fallback(inv_wave, wave)
|
||||
inv_mode = fallback(inv_mode, mode)
|
||||
inv_biort = fallback(inv_biort, biort)
|
||||
inv_qshift = fallback(inv_qshift, qshift)
|
||||
if use_dtcwt:
|
||||
fwdfun, invfun = ptwav.DTCWTForward, ptwav.DTCWTInverse
|
||||
elif use_1d_dwt:
|
||||
fwdfun, invfun = ptwav.DWT1DForward, ptwav.DWT1DInverse
|
||||
else:
|
||||
fwdfun, invfun = ptwav.DWTForward, ptwav.DWTInverse
|
||||
if use_dtcwt:
|
||||
self._wavelet_forward = fwdfun(
|
||||
J=level,
|
||||
mode=mode,
|
||||
biort=biort,
|
||||
qshift=qshift,
|
||||
)
|
||||
self._wavelet_inverse = invfun(
|
||||
mode=inv_mode,
|
||||
biort=inv_biort,
|
||||
qshift=inv_qshift,
|
||||
)
|
||||
else:
|
||||
self._wavelet_forward = fwdfun(J=level, wave=wave, mode=mode)
|
||||
self._wavelet_inverse = invfun(wave=inv_wave, mode=inv_mode)
|
||||
if device is not None:
|
||||
self._wavelet_forward = self._wavelet_forward.to(device=device)
|
||||
self._wavelet_inverse = self._wavelet_inverse.to(device=device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
*,
|
||||
forward_function: Callable | None = None,
|
||||
) -> tuple[torch.Tensor, tuple]:
|
||||
return fallback(forward_function, self._wavelet_forward)(t)
|
||||
|
||||
def inverse(
|
||||
self,
|
||||
yl: torch.Tensor,
|
||||
yh: tuple,
|
||||
*,
|
||||
inverse_function: Callable | None = None,
|
||||
two_step_inverse: bool = False,
|
||||
) -> torch.Tensor:
|
||||
inverse_function = fallback(inverse_function, self._wavelet_inverse)
|
||||
if not two_step_inverse:
|
||||
return inverse_function((yl, yh))
|
||||
result = inverse_function((torch.zeros_like(yl), yh))
|
||||
result += inverse_function((
|
||||
yl,
|
||||
tuple(torch.zeros_like(yh_band) for yh_band in yh),
|
||||
))
|
||||
return result
|
||||
|
||||
def to(self, *args: list, copy: bool = False, **kwargs: dict) -> Wavelet:
|
||||
o = Wavelet.__new__(Wavelet) if copy else self
|
||||
o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001
|
||||
o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def wavelist() -> tuple:
|
||||
return tuple(pywt.wavelist()) if HAVE_WAVELETS else ()
|
||||
|
||||
@staticmethod
|
||||
def biortlist() -> tuple:
|
||||
return (
|
||||
("near_sym_a", "near_sym_b", "antonini", "legall") if HAVE_WAVELETS else ()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def qshiftlist() -> tuple:
|
||||
return (
|
||||
("qshift_a", "qshift_b", "qshift_c", "qshift_d", "qshift_06")
|
||||
if HAVE_WAVELETS
|
||||
else ()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def modelist() -> tuple:
|
||||
return (
|
||||
(
|
||||
"symmetric",
|
||||
"zero",
|
||||
"reflect",
|
||||
"replicate",
|
||||
"periodization",
|
||||
"periodic",
|
||||
"constant",
|
||||
)
|
||||
if HAVE_WAVELETS
|
||||
else ()
|
||||
)
|
||||
|
||||
|
||||
def expand_yh_scales(
|
||||
yh: Sequence,
|
||||
*,
|
||||
yh_scales: float | Sequence = 1.0,
|
||||
) -> float | tuple:
|
||||
yhlen = len(yh)
|
||||
yh_shape = yh[0].shape
|
||||
# Doesn't make sense to target orientations for 1D DWD (3D here).
|
||||
olen = yh_shape[2] if len(yh_shape) > 3 else 1
|
||||
# print(f"\nSIZES: yhlen={yhlen}, olen={olen}, yh_shape={yh[0].shape}")
|
||||
if isinstance(yh_scales, (float, int)):
|
||||
return ((float(yh_scales),) * olen,) * yhlen
|
||||
otemplate = (1.0,) * olen
|
||||
yh_scales = tuple(
|
||||
(float(band),) * olen
|
||||
if isinstance(band, (float, int))
|
||||
else (
|
||||
(
|
||||
*(float(i) for i in band[:olen]),
|
||||
*otemplate[: olen - len(band[:olen])],
|
||||
)
|
||||
if isinstance(band, (tuple, list))
|
||||
else band
|
||||
)
|
||||
for band in yh_scales
|
||||
)
|
||||
if "fill" in yh_scales:
|
||||
fillidx = yh_scales.index("fill")
|
||||
if "fill" in yh_scales[fillidx + 1 :]:
|
||||
raise ValueError("Only one fill allowed.")
|
||||
if fillidx == 0 or len(yh_scales) < 2:
|
||||
raise ValueError(
|
||||
"Invalid fill value, cannot be in the first position or the only item.",
|
||||
)
|
||||
yhslen = len(yh_scales)
|
||||
if yhslen - 1 < yhlen:
|
||||
# Need to pad.
|
||||
fill = (yh_scales[fillidx - 1],) * (yhlen - (len(yh_scales) - 1))
|
||||
yh_scales = (*yh_scales[:fillidx], *fill, *yh_scales[fillidx + 1 :])
|
||||
else:
|
||||
# Just remove the "fill".
|
||||
yh_scales = (*yh_scales[:fillidx], *yh_scales[fillidx + 1 :])
|
||||
return yh_scales[:yhlen]
|
||||
|
||||
|
||||
def wavelet_scaling(
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
yl_scale: float | torch.Tensor,
|
||||
yh_scales: float | Sequence | None,
|
||||
*,
|
||||
in_place: bool = False,
|
||||
) -> tuple:
|
||||
if not in_place:
|
||||
yl = yl.clone()
|
||||
yh = tuple(yhband.clone() for yhband in yh)
|
||||
if yl_scale != 1.0:
|
||||
yl *= yl_scale
|
||||
yh_scales = expand_yh_scales(
|
||||
yh,
|
||||
yh_scales=yh_scales if yh_scales is not None else 1.0,
|
||||
)
|
||||
for hscale, ht in zip(yh_scales, yh):
|
||||
if isinstance(hscale, (int, float)):
|
||||
ht *= hscale # noqa: PLW2901
|
||||
continue
|
||||
for lidx in range(min(ht.shape[2], len(hscale))):
|
||||
ht[:, :, lidx] *= hscale[lidx]
|
||||
return (yl, yh)
|
||||
|
||||
|
||||
def wavelet_blend(
|
||||
a: tuple,
|
||||
b: tuple,
|
||||
*,
|
||||
yl_factor: torch.Tensor | float,
|
||||
blend_function: Callable,
|
||||
yh_factor: torch.Tensor | float | None = None,
|
||||
yh_blend_function: Callable | None = None,
|
||||
) -> tuple:
|
||||
if not isinstance(yl_factor, torch.Tensor):
|
||||
yl_factor = a[0].new_full((1,), yl_factor)
|
||||
if yh_factor is None:
|
||||
yh_factor = yl_factor
|
||||
elif not isinstance(yh_factor, torch.Tensor):
|
||||
yh_factor = a[0].new_full((1,), yh_factor)
|
||||
yh_blend_function = fallback(yh_blend_function, blend_function)
|
||||
return (
|
||||
blend_function(a[0], b[0], yl_factor),
|
||||
tuple(yh_blend_function(ta, tb, yh_factor) for ta, tb in zip(a[1], b[1])),
|
||||
)
|
||||
Reference in New Issue
Block a user