Add SonarNestedNoise node

Add centering and redistribute mode to quantile normalization features
Various cleanups
This commit is contained in:
blepping
2026-08-17 05:15:32 -06:00
parent bba5bf25e9
commit 650467ce97
13 changed files with 632 additions and 122 deletions
+4 -3
View File
@@ -4,9 +4,10 @@ import contextlib
import importlib
import sys
from functools import partial
from typing import TYPE_CHECKING, Callable, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple
if TYPE_CHECKING:
from collections.abc import Callable
from types import ModuleType
@@ -83,7 +84,7 @@ class Integrations:
class SonarIntegrations(Integrations):
def __init__(self, *args: list, **kwargs: dict):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
self.register_integration(
@@ -114,7 +115,7 @@ MODULES = SonarIntegrations()
class IntegratedNode(type):
@staticmethod
def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict:
def wrap_INPUT_TYPES(orig_method: Callable, *args: Any, **kwargs: Any) -> dict:
MODULES.initialize()
return orig_method(*args, **kwargs)
+12 -12
View File
@@ -2,14 +2,14 @@ from __future__ import annotations
import math
import random
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
import torch
from . import utils
if TYPE_CHECKING:
from types import Sequence
from collections.abc import Sequence
class SonarLatentOperation:
@@ -34,9 +34,9 @@ class SonarLatentOperation:
def call_op(
self,
t: torch.Tensor,
*args: list,
*args: Any,
op=None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
if op is None:
op = self.op
@@ -51,7 +51,7 @@ class SonarLatentOperation:
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
if not self.enabled(sigma=sigma):
return latent
@@ -70,7 +70,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
difference_multiplier: float,
ops: Sequence,
op_alt=None,
**kwargs: dict,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self.blend_function = utils.BLENDING_MODES[blend_mode]
@@ -87,7 +87,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
t = latent
enabled = self.enabled(sigma)
@@ -115,13 +115,13 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
class SonarLatentOperationNoise(SonarLatentOperation):
def __init__(
self,
*args: list,
*args: Any,
custom_noise,
scale_to_sigma: bool = False,
cpu_noise: bool = False,
normalize: bool = True,
lazy_noise_sampler: bool = False,
**kwargs: dict,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.custom_noise = custom_noise
@@ -137,7 +137,7 @@ class SonarLatentOperationNoise(SonarLatentOperation):
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
t = latent
enabled = self.enabled(sigma)
@@ -193,12 +193,12 @@ class SonarLatentOperationNoise(SonarLatentOperation):
class SonarLatentOperationSetSeed(SonarLatentOperation):
def __init__(self, *args: list, seed: int, restore_rng_state: bool, **kwargs: dict):
def __init__(self, *args: Any, seed: int, restore_rng_state: bool, **kwargs: Any):
super().__init__(*args, **kwargs)
self.seed = seed
self.restore_rng_state = restore_rng_state
def __call__(self, *args: list, **kwargs: dict) -> torch.Tensor:
def __call__(self, *args: Any, **kwargs: Any) -> torch.Tensor:
if self.restore_rng_state:
pyrandst = random.getstate()
torchrandst = torch.random.get_rng_state()
+3 -3
View File
@@ -25,11 +25,11 @@ NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
if not HAVE_COMFY_UNION_TYPE:
class Wildcard(str): # noqa: FURB189
class Wildcard(str):
__slots__ = ("whitelist",)
@classmethod
def __new__(cls, s, *args: list, whitelist=None, **kwargs: Any):
def __new__(cls, s, *args: Any, whitelist=None, **kwargs: Any):
result = super().__new__(s, *args, **kwargs)
result.whitelist = whitelist
return result
@@ -186,7 +186,7 @@ class SonarInputTypes(InputTypes):
class SonarLazyInputTypes(LazyInputTypes):
_NO_REPLACE = True
def __init__(self, *args: list, initializers=(MODULES.initialize,), **kwargs: Any):
def __init__(self, *args: Any, initializers=(MODULES.initialize,), **kwargs: Any):
super().__init__(
*args,
initializers=initializers,
+59 -50
View File
@@ -3,29 +3,38 @@ from __future__ import annotations
from copy import deepcopy
from functools import partial
from typing import Callable, TypeVar
from typing import TYPE_CHECKING, Any, TypeVar
if TYPE_CHECKING:
from collections.abc import Callable
bi_int = int
bi_bool = bool
bi_float = float
class InputCollection:
_DELEGATE_KEYS = frozenset((
"bool",
"boolean",
"clip",
"conditioning",
"field",
"float",
"image",
"int",
"latent",
"model",
"sampler",
"seed",
"sigmas",
"string",
"vae",
))
_DELEGATE_KEYS = frozenset(
(
"bool",
"boolean",
"clip",
"conditioning",
"field",
"float",
"image",
"int",
"latent",
"model",
"sampler",
"seed",
"sigmas",
"string",
"vae",
),
)
def __init__(self, **kwargs: dict):
def __init__(self, **kwargs: Any):
self.fields = kwargs
def __getattr__(self, key: str):
@@ -42,10 +51,10 @@ class InputCollection:
def clone(self):
return InputCollection(**self.to_dict())
def __len__(self) -> int:
def __len__(self) -> bi_int:
return len(self.fields)
def __contains__(self, key: str) -> bool:
def __contains__(self, key: str) -> bi_bool:
return key in self.fields
def field(
@@ -53,8 +62,8 @@ class InputCollection:
name: str,
type: str | tuple,
*,
_skip: bool = False,
**kwargs: dict,
_skip: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
if not _skip:
self.fields[name] = (type,) if not kwargs else (type, kwargs)
@@ -63,7 +72,7 @@ class InputCollection:
def string(
self,
name: str,
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "STRING", **kwargs)
@@ -71,11 +80,11 @@ class InputCollection:
self,
name: str,
*,
step: float = 0.001,
min: float = -10000.0,
max: float = 10000.0,
round: bool = False,
**kwargs: dict,
step: bi_float = 0.001,
min: bi_float = -10000.0,
max: bi_float = 10000.0,
round: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
return self.field(
name,
@@ -91,9 +100,9 @@ class InputCollection:
self,
name: str,
*,
min: float = -10000,
max: float = 10000,
**kwargs: dict,
min: bi_float = -10000,
max: bi_float = 10000,
**kwargs: Any,
) -> InputCollection:
return self.field(
name,
@@ -106,22 +115,22 @@ class InputCollection:
def bool(
self,
name: str,
default: bool = False,
**kwargs: dict,
default: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "BOOLEAN", default=default, **kwargs)
boolean = bool
boolean = bool # noqa: A003
def seed(
self,
name: str = "seed",
*,
default: int = 0,
min: int = 0,
max: int = 0xFFFFFFFFFFFFFFFF,
default: bi_int = 0,
min: bi_int = 0,
max: bi_int = 0xFFFFFFFFFFFFFFFF,
tooltip="Seed to use for generated noise",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.int(
name,
@@ -132,32 +141,32 @@ class InputCollection:
**kwargs,
)
def image(self, name: str = "image", **kwargs: dict) -> InputCollection:
def image(self, name: str = "image", **kwargs: Any) -> InputCollection:
return self.field(name, "IMAGE", **kwargs)
def latent(self, name: str = "latent", **kwargs: dict) -> InputCollection:
def latent(self, name: str = "latent", **kwargs: Any) -> InputCollection:
return self.field(name, "LATENT", **kwargs)
def conditioning(
self,
name: str = "conditioning",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "CONDITIONING", **kwargs)
def model(self, name: str = "model", **kwargs: dict) -> InputCollection:
def model(self, name: str = "model", **kwargs: Any) -> InputCollection:
return self.field(name, "MODEL", **kwargs)
def sigmas(self, name: str = "sigmas", **kwargs: dict) -> InputCollection:
def sigmas(self, name: str = "sigmas", **kwargs: Any) -> InputCollection:
return self.field(name, "SIGMAS", **kwargs)
def sampler(self, name: str = "sampler", **kwargs: dict) -> InputCollection:
def sampler(self, name: str = "sampler", **kwargs: Any) -> InputCollection:
return self.field(name, "SAMPLER", **kwargs)
def clip(self, name: str = "clip", **kwargs: dict) -> InputCollection:
def clip(self, name: str = "clip", **kwargs: Any) -> InputCollection:
return self.field(name, "CLIP", **kwargs)
def vae(self, name: str = "vae", **kwargs: dict) -> InputCollection:
def vae(self, name: str = "vae", **kwargs: Any) -> InputCollection:
return self.field(name, "VAE", **kwargs)
@@ -226,7 +235,7 @@ class InputTypes:
errstr = f"Unknown attribute {key} for InputTypes"
raise AttributeError(errstr)
def wrapper(*args: list, **kwargs: dict):
def wrapper(*args: Any, **kwargs: Any):
meth(*args, **kwargs)
return self
@@ -240,7 +249,7 @@ class LazyInputTypes:
self.builder = builder
self.initializers = initializers
def get_input_types(self, *args: list, **kwargs: dict):
def get_input_types(self, *args: Any, **kwargs: Any):
if args or kwargs:
args = tuple(args)
cache_key = (args, tuple(kwargs.items()))
@@ -259,5 +268,5 @@ class LazyInputTypes:
self._input_types_params[cache_key] = result
return result
def __call__(self, *args: list, **kwargs: dict) -> dict:
def __call__(self, *args: Any, **kwargs: Any) -> dict:
return self.get_input_types(*args, **kwargs)()
+5
View File
@@ -351,8 +351,10 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
abs_quantiles: bool,
only_outliers: bool,
manual_quantiles: str,
zero_mean_scale: str,
):
# TODO: Support an optional reference LATENT_OPERATION.
zms, rms, zrs = cls._parse_mean_scales(zero_mean_scale)
nq_lo, nq_hi = cls._parse_manual_quantiles(manual_quantiles)
qnorm_filter = functools.partial(
utils.quantile_normalize,
@@ -368,6 +370,9 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
only_outliers=only_outliers,
nq_lo=nq_lo,
nq_hi=nq_hi,
zero_mean_scale=zms,
restore_mean_scale=rms,
zero_result_mean_scale=zrs,
)
return (SonarLatentOperation(op=lambda latent: qnorm_filter(latent)),) # noqa: PLW0108
+33 -10
View File
@@ -30,7 +30,6 @@ if TYPE_CHECKING:
try:
from comfy import nested_tensor
from comfy.utils import pack_latents, unpack_latents
except (ModuleNotFoundError, ImportError):
nested_tensor = None
@@ -444,9 +443,6 @@ class CustomNOISE:
batch_size = latent_image.shape[0]
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
# use_idxs = {
# out_idx: idx % batch_size for out_idx, idx in enumerate(unique_inds)
# }
use_idxs = (idx for idx in range(unique_inds[-1] + 1) if idx in unique_inds)
use_idxs = {idx: inverse_inds[uidx] for uidx, idx in enumerate(use_idxs)}
n_use_idxs = len(use_idxs)
@@ -462,9 +458,6 @@ class CustomNOISE:
sample_idx = idx % batch_size
for nidx in range(len(nested_parts)):
sample = nested_parts[nidx][sample_idx].unsqueeze(0)
# print(
# f"\nNOISE: idx {idx}.{nidx}, sample_idx {sample_idx}, shape {latent_image[sample_idx].shape}, nested={sample.is_nested}",
# )
noise = self._sample_noise(sample, self.seed + idx, sampler_idx=nidx)
batch_out_idx = use_idxs.get(idx)
if batch_out_idx is not None:
@@ -478,7 +471,7 @@ class CustomNOISE:
class SonarToComfyNOISENode(metaclass=IntegratedNode):
DESCRIPTION = "Allows converting SONAR_CUSTOM_NOISE to NOISE (used by SamplerCustomAdvanced and possibly other custom samplers). NOTE: Does not work with noise types that depend on sigma (Brownian, ScheduledNoise, etc)."
DESCRIPTION = "Allows converting SONAR_CUSTOM_NOISE to NOISE (used by SamplerCustomAdvanced and possibly other custom samplers). The extra alt inputs are used if the latent is a nested tensor and ignored otherwise. Audio/video models like LTX and MiniMax H3 used nested tensors (order video then audio). Connected custom noise inputs will be used in order. NOTE: This node does not work with noise types that depend on sigma (Brownian, ScheduledNoise, etc) unless you manually set a sigma via other nodes."
RETURN_TYPES = ("NOISE",)
CATEGORY = "sampling/custom_sampling/noise"
FUNCTION = "go"
@@ -502,14 +495,44 @@ class SonarToComfyNOISENode(metaclass=IntegratedNode):
default=1.0,
tooltip="Simple multiplier applied to noise after all other scaling and normalization effects. If set to 0, no noise will be generated (same as disabling noise).",
)
.opt_customnoise_alt_custom_noise_1(
tooltip="Optional custom noise. See the node description.",
)
.opt_customnoise_alt_custom_noise_2(
tooltip="Optional custom noise. See the node description.",
)
.opt_customnoise_alt_custom_noise_3(
tooltip="Optional custom noise. See the node description.",
)
),
)
@classmethod
def go(cls, *, custom_noise, seed, cpu_noise=True, normalize=True, multiplier=1.0):
def go(
cls,
*,
custom_noise,
seed,
cpu_noise=True,
normalize=True,
multiplier=1.0,
alt_custom_noise_1=None,
alt_custom_noise_2=None,
alt_custom_noise_3=None,
):
noises = tuple(
cn.clone()
for cn in (
custom_noise,
alt_custom_noise_1,
alt_custom_noise_2,
alt_custom_noise_3,
)
if cn is not None
)
return (
CustomNOISE(
(custom_noise,),
noises,
seed,
cpu_noise=cpu_noise,
normalize=normalize,
+89 -3
View File
@@ -1,8 +1,15 @@
from __future__ import annotations
import math
import torch
from comfy import model_management
try:
from comfy import nested_tensor
except (ModuleNotFoundError, ImportError):
nested_tensor = None
from .. import noise, utils
from ..latent_ops import SonarLatentOperation
from ..sonar import SonarGuidanceMixin
@@ -737,7 +744,7 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
tooltip="Multiplier on the input noise just before it is clipped to the quantile min/max. Generally should be left at the default.",
)
.req_float_norm_power(
default=0.5,
default=0.0,
min=-10000.0,
max=10000.0,
step=0.001,
@@ -764,9 +771,9 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
tooltip="Same as norm_power, except this is applied before the quantile is calculated. For example, you could set this to 2.0 and norm_power to 0.5 to do quantile normalization on squared noise then return the square root as the result.",
)
.req_field_sign_mode(
("default", "keep", "avoid"),
("default", "keep", "avoid", "flip", "pos", "neg"),
default="default",
tooltip="This setting will be overridden by strategies that end with keepsign or avoidsign. Controls whether outliers will keep or avoid the sign of the original noise input. If set to default then it depends on whether the strategy affects the sign.",
tooltip="This setting will be overridden by strategies that end with keepsign or avoidsign. Controls whether outliers will keep or avoid the sign of the original noise input. If set to default then it depends on whether the strategy affects the sign. You can also set this to flip, pos or neg to flip or force the sign on outliers without regard for the original value.",
)
.req_bool_abs_quantiles(
default=True,
@@ -780,6 +787,10 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
default="",
tooltip="One or two comma separated float values. Can be used to manually define the quantile *values* (not percentages). When specified, the quantile parameter and noise reference parameters are ignored. The sign is always ignored. When abs_quantiles is enabled, only the first value will be used otherwise the first value defines the negative boundary and the second defines the positive. Example with abs_quantiles enabled: You enter '3.0', the range will be -3.0 to 3.0. Example with abs_quantiles disabled: You enter '2.0, 3.0', the range will be -2.0 to 3.0. As a reference for ranges to set, out of 100,000,000 items of Gaussian noise, about 100,000 will be above 2.6 so using limits between 2.5-3.5 will be roughly in the range of Gaussian noise.",
)
.req_string_zero_mean_scale(
default="0.0",
tooltip="Up to three comma separated float values. The first controls subtracting the mean from the input (and reference if provided), the second controls restoring the mean, the third controls subtracting mean from the quantile-normalized result. This will operate over the same dimensions as quantile normalization. For example, '1, 1, 0' would mean subtract the mean, then restore it, but _don't_ zero the mean before restoration. When not provided, the second two values will use the same parameter as the first. Leaving the field empty is the same as setting it to 0.",
)
.opt_customnoise_reference_noise_opt(
tooltip="If connected then noise from this generator will be used to calculate the quantiles but normalization will be applied to the custom_noise input. This probably won't work correctly for negative quantiles. When used, the reference will be generated first so, for example, you could use the SonarCustomNoiseParameters in forked RNG node to have both generate noise with the same seed/state. Ignored when specifying manual quantiles.",
)
@@ -807,6 +818,19 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
nq_hi = split_quantiles[0 if len(split_quantiles) < 2 else 1]
return nq_lo, nq_hi
@staticmethod
def _parse_mean_scales(zero_mean_scale: str) -> tuple[float, float, float]:
zero_mean_scales = tuple(
0.0 if not s.strip() else float(s)
for s in zero_mean_scale.strip().split(",")
)
zmlen = len(zero_mean_scales)
if zmlen > 3:
raise ValueError(
"Too many values provided for zero mean scales. Must be between 0 and 3.",
)
return tuple(zero_mean_scales[idx if idx < zmlen else 0] for idx in range(3))
def go(
self,
*,
@@ -825,8 +849,10 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
abs_quantiles: bool,
only_outliers: bool,
manual_quantiles: str,
zero_mean_scale: str,
reference_noise_opt: object | None = None,
) -> tuple:
zms, rms, zrs = self._parse_mean_scales(zero_mean_scale)
nq_lo, nq_hi = self._parse_manual_quantiles(manual_quantiles)
return super().go(
factor,
@@ -846,6 +872,9 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase):
only_outliers=only_outliers,
nq_lo=nq_lo,
nq_hi=nq_hi,
zero_mean_scale=zms,
restore_mean_scale=rms,
zero_result_mean_scale=zrs,
)
@@ -1728,6 +1757,62 @@ class SonarCustomNoiseParametersNode(
)
class SonarNestedNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "Custom noise helper that allows working with nested tensors which are generally only used by audio/video models like MiniMax H3, LTX, etc. Those models use order video then audio. The custom noise inputs are optional but at least one must be attached and will be used in order. You may supply less custom noise inputs than nested parts, in which case the noise used will wrap. In other words, if there were three nested components and you attach two custom noise inputs, the third component will use the first custom noise item. Note: Aside from batch, the number of elements in the latent reference you supply and the latent you use for sampling must match."
INPUT_TYPES = SonarLazyInputTypes(
lambda: (
NoiseNoChainInputTypes()
.req_latent_latent(
tooltip="Latent to use as a reference for nested shapes.",
)
.opt_customnoise_custom_noise_1()
.opt_customnoise_custom_noise_2()
.opt_customnoise_custom_noise_3()
.opt_customnoise_custom_noise_4()
),
)
@classmethod
def get_item_class(cls):
return noise.NestedNoise
def go(
self,
*,
factor: float,
latent: dict,
custom_noise_1=None,
custom_noise_2=None,
custom_noise_3=None,
custom_noise_4=None,
):
noises = tuple(
cn
for cn in (custom_noise_1, custom_noise_2, custom_noise_3, custom_noise_4)
if cn is not None
)
if not noises:
raise ValueError("At least one custom noise item must be attached.")
samples = latent["samples"]
if nested_tensor is not None and samples.is_nested:
nested_parts = samples.unbind()
else:
nested_parts = (samples,)
n_parts = len(nested_parts)
if not n_parts:
raise ValueError("Latent reference apparently has 0 parts?")
noises = tuple(cn.clone() for cn in noises[:n_parts])
shapes = tuple((1, *t.shape[1:]) for t in nested_parts)
numel = sum(math.prod(shp) for shp in shapes)
return super().go(
factor,
custom_noises=noises,
nested_numel=numel,
nested_shapes=shapes,
)
NODE_CLASS_MAPPINGS = {
"SonarBlendedNoise": SonarBlendedNoiseNode,
"SonarChannelNoise": SonarChannelNoiseNode,
@@ -1736,6 +1821,7 @@ NODE_CLASS_MAPPINGS = {
"SonarGuidedNoise": SonarGuidedNoiseNode,
"SonarLatentOperationFilteredNoise": SonarLatentOperationFilteredNoiseNode,
"SonarModulatedNoise": SonarModulatedNoiseNode,
"SonarNestedNoise": SonarNestedNoiseNode,
"SonarNormalizeNoiseToScale": SonarNormalizeNoiseToScaleNode,
"SonarNoveltyFilteredNoise": SonarNoveltyFilteredNoiseNode,
"SonarPatternBreakNoise": SonarPatternBreakNoiseNode,
+7 -6
View File
@@ -7,6 +7,7 @@ from __future__ import annotations
import math
import os
import random
from typing import Any
import comfy
import folder_paths
@@ -86,7 +87,7 @@ class ChannelMixer:
channel_mixer /= channel_mixer.norm(dim=1, keepdim=True)
return channel_mixer
def to(self, *args: list, **kwargs: dict):
def to(self, *args: Any, **kwargs: Any):
if self.mixer is not None:
self.mixer = self.mixer.to(*args, **kwargs)
return self
@@ -100,7 +101,7 @@ class ChannelMixer:
noise = self.mixer @ noise.swapaxes(0, 1).reshape(c, -1)
return noise.reshape(c, b, h, w).swapaxes(1, 0)
def __call__(self, *args: list, **kwargs: dict):
def __call__(self, *args: Any, **kwargs: Any):
return self.apply(*args, **kwargs)
@@ -301,7 +302,7 @@ class PowerNoiseItem(CustomNoiseItemBase):
*,
channel_correlation,
power_filter=None,
**kwargs: dict,
**kwargs: Any,
):
if isinstance(channel_correlation, str):
channel_correlation = torch.tensor(
@@ -476,7 +477,7 @@ class PowerFilterNoiseItem(PowerNoiseItem):
noise,
normalize_noise,
normalize_result,
**kwargs: dict,
**kwargs: Any,
):
super().__init__(
factor,
@@ -628,7 +629,7 @@ class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
def go(
self,
preview="none",
**kwargs: dict,
**kwargs: Any,
):
result = super().go(**kwargs)
if preview == "none":
@@ -713,7 +714,7 @@ class SonarPowerFilterNoiseNode(SonarPowerNoiseNode, SonarNormalizeNoiseNodeMixi
normalize_noise,
normalize_result,
preview="none",
**kwargs: dict,
**kwargs: Any,
):
return super().go(
factor=factor,
+68
View File
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING
import comfy
import torch
import yaml
from comfy import utils as comfy_utils
from comfy.k_diffusion import sampling
from comfy.model_management import throw_exception_if_processing_interrupted
from torch import Tensor
@@ -25,6 +26,11 @@ from .utils import (
scale_noise,
)
try:
from comfy import nested_tensor
except (ModuleNotFoundError, ImportError):
nested_tensor = None
if TYPE_CHECKING:
from collections.abc import Callable
@@ -1929,6 +1935,9 @@ class QuantileFilteredNoise(CustomNoiseItemBase):
only_outliers=self.only_outliers,
nq_lo=self.nq_lo,
nq_hi=self.nq_hi,
zero_mean_scale=self.zero_mean_scale,
restore_mean_scale=self.restore_mean_scale,
zero_result_mean_scale=self.zero_result_mean_scale,
)
def noise_sampler(*args, **kwargs):
@@ -2319,6 +2328,65 @@ class CustomNoiseParametersNoise(CustomNoiseItemBase):
return noise_sampler
class NestedNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "custom_noises":
return tuple(n.clone() for n in self.custom_noises)
return super().clone_key(k)
def make_noise_sampler(
self,
x,
sigma_min,
sigma_max,
*args,
**kwargs,
):
x_numel = x[0].numel()
x_shape = x.shape
shapes = self.nested_shapes
if x_numel != self.nested_numel:
errstr = f"Sampling latent with shape {x.shape} has {x_numel} element(s) per batch item which does not match the initial reference with shape(s) {shapes} and element(s) {self.nested_numel}"
raise ValueError(errstr)
factor = self.factor
n_shapes = len(shapes)
# batch_size = x.shape[0]
noise_samplers = tuple(
n.make_noise_sampler(
x_part,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
**kwargs,
)
for n, x_part in zip(
self.custom_noises[:n_shapes],
comfy_utils.unpack_latents(x, shapes),
strict=True,
)
)
del x
n_samplers = len(noise_samplers)
if not n_samplers:
raise RuntimeError("Internal error: No noise samplers available")
def noise_sampler(*args, **kwargs) -> torch.Tensor:
noise_chunks = (
noise_samplers[nidx % n_samplers](*args, **kwargs)
for nidx in range(n_shapes)
)
if factor != 1:
noise_chunks = (nc.mul_(factor) for nc in noise_chunks)
result = comfy_utils.pack_latents(tuple(noise_chunks))[0]
return (
result
if result.shape[1:] == x_shape[1:]
else result.reshape(-1, *x_shape[1:])
)
return noise_sampler
class BlehOpsNoise(CustomNoiseItemBase):
def __init__(
self,
@@ -1,7 +1,7 @@
# Some noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers
from __future__ import annotations
from typing import TYPE_CHECKING, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
@@ -11,8 +11,6 @@ from .base import FramesToChannelsNoiseGenerator
if TYPE_CHECKING:
from collections.abc import Sequence
# ruff: noqa: ANN002, ANN003
class WaveletNoiseOctave(NamedTuple):
octave: int
@@ -47,7 +45,7 @@ class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
"noise_sampler": None,
}
def __init__(self, *args, **kwargs):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.set_octave_data()
@@ -96,7 +94,7 @@ class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
raise ValueError("Unworkable parameters for wavelet noise")
self.octave_data = tuple(octave_data)
def _generate_octave(self, *args: list, shape: Sequence) -> torch.Tensor:
def _generate_octave(self, *args: Any, shape: Sequence) -> torch.Tensor:
height, width = shape[-2:]
noise = (
self.noise_sampler(*args)[..., :height, :width].reshape(shape)
@@ -122,7 +120,7 @@ class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
self.update_blend,
)
def generate(self, *args: list) -> torch.Tensor:
def generate(self, *args: Any) -> torch.Tensor:
adjusted_shape = self.get_adjusted_shape()
height, width = adjusted_shape[-2:]
curr_shape = list(adjusted_shape)
+13 -14
View File
@@ -247,7 +247,7 @@ class SonarBase:
momentum = self.cfg.momentum if momentum is None else momentum
mode = self.cfg.momentum_mode
if (
momentum == 1 # noqa: PLR0916
momentum == 1
or history is None
or (mode == MomentumMode.DENOISED and not is_denoised)
or (mode != MomentumMode.DENOISED and is_denoised)
@@ -412,7 +412,7 @@ class SonarGuidanceMixin:
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
def __init__(self, *args: list[Any], **kwargs: dict[str, Any]):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
SonarGuidanceMixin.__init__(self, self.cfg.guidance)
@@ -424,8 +424,8 @@ class SonarSampler(SonarWithGuidance):
sigmas: Tensor,
s_in: Tensor,
extra_args: dict[str, Any],
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.model = model
@@ -437,7 +437,7 @@ class SonarSampler(SonarWithGuidance):
self,
x: Tensor,
sigma: Tensor,
*args: list[Any],
*args: Any,
s_in=None,
extra_args=None,
) -> Tensor:
@@ -452,8 +452,8 @@ class SonarSampler(SonarWithGuidance):
class SonarEuler(SonarSampler):
def __init__(
self,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
@@ -531,8 +531,8 @@ class SonarEulerAncestral(SonarSampler):
self,
eta: float = 1.0,
s_noise: float = 1.0,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.eta = eta
@@ -560,9 +560,8 @@ class SonarEulerAncestral(SonarSampler):
)
if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
result_sample = ( # noqa: PLR6104
result_sample
+ self.noise_sampler(sigma, sigma_next) * (self.s_noise * sigma_up)
result_sample = result_sample + self.noise_sampler(sigma, sigma_next) * (
self.s_noise * sigma_up
)
return (
@@ -630,8 +629,8 @@ class SonarDPMPPSDE(SonarSampler):
self,
eta: float = 1.0,
s_noise: float = 1.0,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.eta = eta
+326 -8
View File
@@ -255,6 +255,55 @@ def _quantile_norm_replace(
return torch.where(mask, noise, candidates)
def _quantile_norm_redistribute(
noise: torch.Tensor,
nq: torch.Tensor,
*,
headroom_multiplier: float = 1.0,
headroom_multiplier_neg: float | None = None,
eps: float | None = None,
) -> torch.Tensor:
if eps is None:
eps = torch.finfo(noise.dtype).eps * 1.25
nq_neg = -nq
t_clamped = torch.clamp(noise, nq_neg, nq)
t_signs = noise.signbit()
excess = noise - t_clamped
headroom = torch.where(
t_signs,
nq_neg - t_clamped,
nq - t_clamped,
).masked_fill_(excess != 0, 0.0)
total_headroom_pos = (
headroom.clamp_min(0)
.sum(dim=-1, keepdim=True)
.mul_(headroom_multiplier)
.clamp_min_(eps)
)
total_headroom_neg = (
headroom.clamp_max(0)
.sum(dim=-1, keepdim=True)
.mul_(
headroom_multiplier_neg
if headroom_multiplier_neg is not None
else headroom_multiplier,
)
.clamp_max_(-eps)
)
factor_pos = (
excess.clamp_min(0.0).sum(dim=-1, keepdim=True).div_(total_headroom_pos)
)
factor_neg = (
excess.clamp_max_(0.0).sum(dim=-1, keepdim=True).div_(total_headroom_neg)
)
t_clamped += torch.where(t_signs, headroom * factor_neg, headroom * factor_pos)
return t_clamped
quantile_handlers = {
"clamp": lambda noise, nq, **_kwargs: noise.clamp(-nq, nq),
"scale_down": _quantile_norm_scaledown,
@@ -444,11 +493,287 @@ quantile_handlers = {
nq,
stiffness=5.0,
),
"redistribute": lambda noise, nq, **_kwargs: _quantile_norm_redistribute(noise, nq),
"redistribute_hr075": lambda noise, nq, **_kwargs: _quantile_norm_redistribute(
noise,
nq,
headroom_multiplier=0.75,
),
"redistribute_hr150": lambda noise, nq, **_kwargs: _quantile_norm_redistribute(
noise,
nq,
headroom_multiplier=1.5,
),
}
def flip_tensor_range(
x: torch.Tensor,
*,
min_neg: torch.Tensor | None = None,
max_pos: torch.Tensor | None = None,
return_ranges: bool = False,
dim: int = -1,
eps: float | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
if eps is None:
eps = torch.finfo(x.dtype).eps * 1.25
# 1. Use the provided maximum positive values, or calculate them dynamically
if max_pos is None:
max_pos = (
torch.clamp_min(x, 0.0).max(dim=dim, keepdim=True).values.clamp_min_(eps)
)
# 2. Use the provided minimum negative values, or calculate them dynamically
if min_neg is None:
min_neg = (
torch.clamp_max(x, 0.0).min(dim=dim, keepdim=True).values.clamp_max_(-eps)
)
# 3. Separate positive and negative elements
is_pos = x >= 0
# 4. Flip positive side: [0, max_pos] -> [eps, max_pos + eps]
x_pos = x.clamp_min(eps)
flipped_pos = (max_pos + eps) - x_pos
# 5. Flip negative side: [min_neg, 0] -> [min_neg - eps, -eps]
x_neg = x.clamp_max(-eps)
flipped_neg = (min_neg - eps) - x_neg
# 6. Recombine the domains
result = torch.where(is_pos, flipped_pos, flipped_neg)
return (result, max_pos, min_neg) if return_ranges else result
# Initial version based on StudentT distribution normalization from https://github.com/Clybius/ComfyUI-Extra-Samplers/
def quantile_normalize(
noise: torch.Tensor,
*,
noise_reference: torch.Tensor | None = None,
quantile: float | tuple | list = 0.75,
dim: int | None = 1,
flatten: bool = True,
nq_fac: float = 1.0,
pow_fac: float = 0.5,
pow_fac_in: float = 0.0,
strategy: str = "clamp",
strategy_handler=None,
# None, keep, avoid
sign_mode: str | None = None,
only_outliers: bool = False,
abs_quantiles: bool = True,
nq_lo: float | None = None,
nq_hi: float | None = None,
zero_mean_scale: float = 0.0,
restore_mean_scale: float | None = None,
zero_result_mean_scale: float | None = None,
eps: float | None = None,
) -> torch.Tensor:
if zero_result_mean_scale is None:
zero_result_mean_scale = zero_mean_scale
if restore_mean_scale is None:
restore_mean_scale = zero_mean_scale
if noise.numel() == 0:
return noise
if dim is not None and dim < 0:
dim = noise.ndim + dim
eff_dim = dim if dim is not None else 0
if eps is None:
eps = torch.finfo(noise.dtype).eps * 1.25
while len(stratparts := strategy.rsplit("_", 1)) == 2 and stratparts[-1] in {
"keepsign",
"avoidsign",
"outliers",
}:
strategy, stratadjust = stratparts
if stratadjust == "outliers":
only_outliers = True
elif stratadjust in {"keepsign", "avoidsign"}:
sign_mode = "keep" if stratadjust == "keepsign" else "avoid"
if nq_lo is None:
if isinstance(quantile, (tuple, list)):
for q in quantile:
noise = quantile_normalize(
noise=noise,
noise_reference=noise_reference,
quantile=q,
dim=dim,
flatten=flatten,
nq_fac=nq_fac,
pow_fac=pow_fac,
strategy=strategy,
strategy_handler=strategy_handler,
sign_mode=sign_mode,
only_outliers=only_outliers,
abs_quantiles=abs_quantiles,
eps=eps,
)
return noise
if quantile is None or quantile >= 1 or quantile <= -1 or quantile == 0:
return noise
inverted = quantile < 0
nq_pos = nq_neg = None
absquantile = abs(quantile)
orig_shape = noise.shape
if noise.ndim > 1 and flatten:
flatnoise = noise.flatten(start_dim=eff_dim)
eff_dim = -1
else:
flatten = False
flatnoise = noise
if nq_lo is not None:
inverted = False
nq_neg = noise.new_tensor(max(eps, abs(nq_lo))).reshape((1,) * flatnoise.ndim)
nq_pos = (
nq_neg
if nq_hi is None
else noise.new_tensor(max(eps, abs(nq_hi))).reshape(nq_neg.shape)
)
if zero_mean_scale or restore_mean_scale:
saved_mean = flatnoise.mean(dim=eff_dim, keepdim=True)
if zero_mean_scale:
flatnoise = flatnoise - (
saved_mean * zero_mean_scale if zero_mean_scale != 1 else saved_mean
)
if restore_mean_scale == 0:
saved_mean = None
elif restore_mean_scale != 1:
saved_mean *= restore_mean_scale
else:
saved_mean = None
orig_noise_flat = flatnoise
if inverted:
flatnoise, max_pos, min_neg = flip_tensor_range(
flatnoise,
dim=eff_dim,
return_ranges=True,
)
if noise_reference is None:
noise_reference = flatnoise
elif noise_reference.numel() != flatnoise.numel():
raise ValueError(
"noise_reference must have the same number of elements as noise",
)
else:
noise_reference = noise_reference.to(flatnoise).reshape(flatnoise.shape)
if zero_mean_scale:
noise_reference = noise_reference - noise_reference.mean(
dim=eff_dim,
keepdim=True,
).mul_(zero_mean_scale)
if inverted:
noise_reference: torch.Tensor = flip_tensor_range(
noise_reference,
dim=eff_dim,
return_ranges=False,
)
if pow_fac_in not in {0, 1}:
noise_reference = (
noise_reference.abs().pow_(pow_fac_in).copysign_(noise_reference)
)
handler = (
quantile_handlers.get(strategy)
if strategy_handler is None
else strategy_handler
)
if handler is None:
raise ValueError("Unknown strategy")
handler = partial(handler, orig_noise=noise, dim=dim, flatten=flatten)
need_outliers = only_outliers or sign_mode is not None
outliers_mask = None
if abs_quantiles:
if nq_pos is None:
nq = torch.quantile(
noise_reference.abs(),
absquantile,
dim=eff_dim,
keepdim=True,
)
nq = nq.mul_(nq_fac).add_(eps)
else:
nq = nq_pos
if need_outliers:
outliers_mask = flatnoise < -nq
outliers_mask |= flatnoise > nq
noise = handler(
flatnoise,
nq,
orig_noise=noise,
dim=dim,
flatten=flatten,
)
else:
noise_signs = noise_reference.signbit()
if nq_pos is None or nq_neg is None:
nq_pos = (
torch.nanquantile(
torch.where(noise_signs, torch.nan, noise_reference),
absquantile,
dim=eff_dim,
keepdim=True,
)
.mul_(nq_fac)
.add_(eps)
)
nq_neg = (
torch.nanquantile(
torch.where(noise_signs, noise_reference.abs(), torch.nan),
absquantile,
dim=eff_dim,
keepdim=True,
)
.mul_(nq_fac)
.add_(eps)
)
noise = torch.where(
noise_signs,
handler(flatnoise.neg(), nq_neg).neg_(),
handler(flatnoise, nq_pos),
)
if need_outliers:
outliers_mask = flatnoise < -nq_neg
outliers_mask |= flatnoise > nq_pos
if pow_fac not in {0.0, 1.0}:
noise = noise.abs().pow_(pow_fac).copysign(noise)
if inverted:
noise = flip_tensor_range(
noise,
dim=eff_dim,
min_neg=min_neg,
max_pos=max_pos,
)
if outliers_mask is not None:
noutliers = noise[outliers_mask].clone() if sign_mode else None
if sign_mode in {"keep", "avoid"}:
noutliers = noutliers.copysign_(
(orig_noise_flat if sign_mode == "keep" else orig_noise_flat.neg())[
outliers_mask
],
)
elif sign_mode == "flip":
noutliers = noutliers.neg_()
elif sign_mode in {"pos", "neg"}:
noutliers = noutliers.abs_()
if sign_mode == "neg":
noutliers = noutliers.neg_()
if noutliers is not None:
noise[outliers_mask] = noutliers
if only_outliers:
inv_outliers_mask = ~outliers_mask
noise[inv_outliers_mask] = orig_noise_flat[inv_outliers_mask]
if zero_result_mean_scale:
noise_mean = noise.mean(dim=eff_dim, keepdim=True).mul_(-zero_result_mean_scale)
saved_mean = noise_mean if saved_mean is None else saved_mean.add_(noise_mean)
if saved_mean is not None:
noise += saved_mean
return noise if noise.shape == orig_shape else noise.reshape(orig_shape)
def quantile_normalize_old(
noise: torch.Tensor,
*,
noise_reference: torch.Tensor | None = None,
@@ -506,7 +831,7 @@ def quantile_normalize(
nq_pos = nq_neg = None
orig_shape = noise.shape
if noise.ndim > 1 and flatten:
flatnoise = noise.flatten(start_dim=dim)
flatnoise = noise.flatten(start_dim=dim or 0)
else:
flatten = False
flatnoise = noise
@@ -631,13 +956,6 @@ def quantile_normalize(
return noise if noise.shape == orig_shape else noise.reshape(orig_shape)
# class QuantileNormMode(Enum):
# # Quantile is applied to absolute values.
# SYMMETRIC = auto()
# # Quantile is applied to signed values.
# SEPERATE = auto()
class QuantileNormQuantileMode(Enum):
QUANTILE = auto()
# User supplied value to use as the quantile value.
+9 -7
View File
@@ -1,13 +1,13 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable
from typing import TYPE_CHECKING, Any
import torch
from .utils import fallback
if TYPE_CHECKING:
from collections.abc import Sequence
from collections.abc import Callable, Sequence
try:
import pytorch_wavelets as ptwav
@@ -98,13 +98,15 @@ class Wavelet:
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),
))
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:
def to(self, *args: Any, copy: bool = False, **kwargs: Any) -> 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