Add SonarNestedNoise node
Add centering and redistribute mode to quantile normalization features Various cleanups
This commit is contained in:
+4
-3
@@ -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
@@ -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
@@ -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
@@ -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)()
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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 +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
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user