Added SonarRippleFilteredNoise node.

Added `SonarApplyLatentOperationCFG` node, similar to the built-in `ApplyLatentOperationCFG` node with scheduling and a lot of different application modes.
Added a `SonarLatentOperationQuantileFilter` node that can be used to apply the quantile normalization functioen to the latent during sampling.
A bunch more quantile normalization modes.
Fixed broken quantile normalization dimension handling. Unfortunately this will likely change seeds.
This commit is contained in:
blepping
2025-06-27 14:55:02 -06:00
parent 2b2a76bcbe
commit c2a93d55cb
6 changed files with 857 additions and 66 deletions
+13
View File
@@ -1,6 +1,19 @@
import sys
from . import py # noqa: F401
from .py import freeu_extreme, nodes, powernoise, sonar
def blep_init():
bi = sys.modules.get("_blepping_integrations", {})
if "sonar" in bi:
return
bi["sonar"] = sys.modules[__name__]
sys.modules["_blepping_integrations"] = bi
sonar.add_samplers()
blep_init()
NODE_CLASS_MAPPINGS = nodes.NODE_CLASS_MAPPINGS | {
"SonarPowerNoise": powernoise.SonarPowerNoiseNode,
+9
View File
@@ -2,6 +2,15 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20250627
* Added `SonarRippleFilteredNoise` node.
* Added `SonarApplyLatentOperationCFG` node, similar to the built-in `ApplyLatentOperationCFG` node with scheduling and a lot of different application modes.
* Added a `SonarLatentOperationQuantileFilter` node that can be used to apply the quantile normalization functioen to the latent during sampling.
* A bunch more quantile normalization modes.
* Fixed broken quantile normalization dimension handling. Unfortunately this will likely change seeds.
## 20250612
* Reimplemented Collatz noise with many new features. Unfortunately this breaks existing workflows. If anyone misses the old version, let me know and I can add it back in (might do that anyway).
+6 -3
View File
@@ -32,8 +32,11 @@ class Integrations:
return self.modules.get(key)
@staticmethod
def get_custom_node(name: str) -> ModuleType | None:
module_key = f"custom_nodes.{name}"
def get_custom_node(module_name: str, key: str) -> ModuleType | None:
bi_module = sys.modules.get("_blepping_integrations", {}).get(key)
if bi_module is not None:
return bi_module
module_key = f"custom_nodes.{module_name}"
with contextlib.suppress(StopIteration):
spec = importlib.util.find_spec(module_key)
if spec is None:
@@ -67,7 +70,7 @@ class Integrations:
return
self.initialized = True
for ih in self.handlers:
module = self.get_custom_node(ih.module_name)
module = self.get_custom_node(ih.module_name, ih.key)
if module is None:
continue
if ih.handler is not None:
+541 -18
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import abc
import functools
import inspect
import math
import random
@@ -2502,6 +2503,522 @@ class SonarWaveletNoiseNode(
)
class SonarRippleFilteredNoiseNode(
SonarCustomNoiseNodeBase,
SonarNormalizeNoiseNodeMixin,
):
DESCRIPTION = (
"Custom noise filter that allows applying scaling based on a wave (sin or cos)."
)
@classmethod
def INPUT_TYPES(cls):
result = super().INPUT_TYPES()
result["required"] |= {
"custom_noise": (
WILDCARD_NOISE,
{
"tooltip": f"Custom noise input. \n{NOISE_INPUT_TYPES_HINT}",
},
),
"mode": (
("sin", "cos", "sin_copysign", "cos_copysign"),
{
"default": "cos",
"tooltip": "Function to use for rippling. The copysign variations are not recommended, they will force the noise to the sign of the wave (whether it's above or below the midline) which has an extremely strong effect. If you want to try it, use something like a 1:16 ratio or higher with normal noise.",
},
),
"dim": (
"INT",
{
"default": -1,
"min": -100,
"max": 100,
"tooltip": "Dimension to use for the ripple effect. Negative dimensions count from the end where -1 is the last dimension.",
},
),
"flatten": (
"BOOLEAN",
{
"default": False,
"tooltip": "When enabled, the noise will be flattened starting from (and including) the specified dimension.",
},
),
"offset": (
"FLOAT",
{
"default": 0.0,
"min": -10000,
"max": 10000.0,
"tooltip": "Simple addition to the base value used for the wave.",
},
),
"roll": (
"FLOAT",
{
"default": 0.0,
"min": -10000,
"max": 10000.0,
"tooltip": "Rolls the wave by this many elements each time the noise generator is called. Negative values roll backward.",
},
),
"amplitude_high": (
"FLOAT",
{
"default": 0.25,
"min": -10000,
"max": 10000.0,
"tooltip": "Scale for noise at the highest point of the wave. This adds to the base value (respecting sign). For example, if set to 0.25 you will get noise * 1.25 at that point. It's also possible to use negative values, -0.25 will result in noise * -1.25.",
},
),
"amplitude_low": (
"FLOAT",
{
"default": 0.15,
"min": -10000,
"max": 10000.0,
"tooltip": "Scale for noise at the lowest point of the wave. This subtracts from the base value (respecting sign). For example, if set to 0.25 you will get noise * 0.75 at that point. It's also possible to use negative values, -0.25 will result in noise * -0.75.",
},
),
"period": (
"FLOAT",
{
"default": 3.0,
"min": -10000,
"max": 10000.0,
"tooltip": "Number of oscillations along the specified dimension.",
},
),
"normalize_noise": (
"BOOLEAN",
{
"default": False,
"tooltip": "Controls whether the noise source is normalized before wavelet filtering occurs.",
},
),
"normalize": (
("default", "forced", "disabled"),
{
"tooltip": "Controls whether the generated noise is normalized to 1.0 strength. For weird blend modes, you may want to set this to forced.",
},
),
}
return result
@classmethod
def get_item_class(cls):
return noise.RippleFilteredNoise
def go(
self,
*,
factor,
rescale,
normalize,
mode: str,
dim: int,
flatten: bool,
offset: float,
amplitude_high: float,
amplitude_low: float,
period: float,
roll: float,
normalize_noise: bool,
custom_noise=None,
sonar_custom_noise_opt=None,
):
return super().go(
factor,
rescale=rescale,
sonar_custom_noise_opt=sonar_custom_noise_opt,
mode=mode,
dim=dim,
flatten=flatten,
offset=offset,
amplitude_high=amplitude_high,
amplitude_low=amplitude_low,
period=period,
roll=roll,
normalize=self.get_normalize(normalize),
normalize_noise=normalize_noise,
noise=custom_noise,
)
class SonarApplyLatentOperationCFG(metaclass=IntegratedNode):
DESCRIPTION = "Allows applying a LATENT_OPERATION during sampling. ComfyUI has a few that are builtin and this node pack also includes: SonarLatentOperationQuantileFilter."
RETURN_TYPES = ("MODEL",)
CATEGORY = "latent/advanced/operations"
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"mode": (
(
"cond_sub_uncond",
"denoised_sub_uncond",
"uncond_sub_cond",
"denoised",
"cond",
"uncond",
"model_input",
),
{
"default": "cond_sub_uncond",
"tooltip": "cond_sub_uncond is what ComfyUI's latent operations use. The non-sub_uncond modes likely won't work with pred_flip mode enabled. If you have anything but the denoised options selected, this will use pre-CFG, otherwise it will use post-CFG (unless you are using model_input).",
},
),
"pred_flip_mode": (
"BOOLEAN",
{
"default": False,
"tooltip": "Lets you try to apply the latent operation to the noise prediction rather than the image prediction. Doesn't work properly with the non-sub_uncond modes. No real reason it should be better, just something you can try. Note: The noise prediction gets scaled by the sigma first, in case that's useful information.",
},
),
"require_uncond": (
"BOOLEAN",
{
"default": False,
"tooltip": "When enabled, the operation will be skipped if uncond is unavailable. This will also happen if you choose a mode that requires uncond.",
},
),
"start_sigma": (
"FLOAT",
{
"default": -1.0,
"min": -1.0,
"max": 9999.0,
"tooltip": "Sigma when the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
},
),
"end_sigma": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 9999.0,
},
),
"blend_mode": (
tuple(utils.BLENDING_MODES.keys()),
{
"default": "lerp",
"tooltip": "Controls how the output of the latent operation is blended with the original result.",
},
),
"blend_strength": (
"FLOAT",
{
"default": 0.5,
"step": 0.001,
"min": -1000.0,
"max": 1000.0,
"round": False,
"tooltip": "Strength of the blend. For a normal blend mode like LERP, 1.0 means use 100% of the output from the latent operation, 0.0 means use none of it and only the original value. Note: Blending is applied to the final result of the operations, in other words operation_2 sees a full unblended result from operation_1.",
},
),
"blend_scale_mode": (
(
"none",
"reverse_sampling",
"sampling",
"reverse_enabled_range",
"enabled_range",
"sampling_sin",
"enabled_range_sin",
),
{
"default": "reverse_sampling",
"tooltip": "Can be used to scale the blend strength over time. Basically works like blend_strength * scale_factor (see below)\nnone: Just uses the blend_strength you have set.\nreverse_sampling: The opposite of the model sampling percent, so if you're making a new generation, the beginning of sampling will be 1.0 and the end will be 0.0. The recommended option as applying these operations usually works better toward the beginning of sampling.\nsampling: Same as reverse_sampling, except the beginning will be 0.0 and the end will be 1.0.\nreverse_enabled_range: Flipped percentage of the range between start_sigma and end_sigma.\nenabled_range: Percentage of the range between start_sigma and end_sigma.\nsampling_sin: Uses the sampling percentage with the sine function such that blend_strength will hit the peak value in the middle of the range.\nenabled_range_sin: Similar to sampling_sin except it applies to the percentage of the enabled range.",
},
),
"blend_scale_offset": (
"FLOAT",
{
"default": 0.0,
"min": -1.0,
"max": 1.0,
"tooltip": "Only applies when blend_scale_mode is not none. Adds the offset to the calculated percentage and then clamps it to be between blend_scale_min and blend_scale_max.",
},
),
"blend_scale_min": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1.0,
"tooltip": "Only applies when blend_scale_mode is not none. Minimum value for the blend scale percentage.",
},
),
"blend_scale_max": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"tooltip": "Only applies when blend_scale_mode is not none. Maximum value for the blend scale percentage.",
},
),
},
"optional": {
"operation_1": (
"LATENT_OPERATION",
{
"tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.",
},
),
"operation_2": (
"LATENT_OPERATION",
{
"tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.",
},
),
"operation_3": (
"LATENT_OPERATION",
{
"tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.",
},
),
"operation_4": (
"LATENT_OPERATION",
{
"tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.",
},
),
"operation_5": (
"LATENT_OPERATION",
{
"tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.",
},
),
},
}
@staticmethod
def get_blend_scaling(
*,
model_sampling: object,
scale_mode: str,
sigma: float,
sigma_t_max: torch.Tensor,
start_sigma: float,
end_sigma: float,
offset: float,
min_pct: float,
max_pct: float,
) -> float | torch.Tensor:
if scale_mode == "none":
return 1.0
if scale_mode in {"sampling", "sampling_sin", "reverse_sampling"}:
rev_sampling_pct = (
(model_sampling.timestep(sigma_t_max) / 999).clamp(0, 1).detach().item()
)
result = (
1.0 - rev_sampling_pct if scale_mode == "sampling" else rev_sampling_pct
)
elif scale_mode in {
"enabled_range",
"enabled_range_sin",
"reverse_enabled_range",
}:
rev_range_pct = (sigma - end_sigma) / (start_sigma - end_sigma)
result = (
1.0 - rev_range_pct if scale_mode == "enabled_range" else rev_range_pct
)
else:
raise ValueError("Bad blend_scale_mode")
if scale_mode.endswith("_sin"):
result = math.sin(result * math.pi)
return max(min_pct, min(result + offset, max_pct))
@classmethod
def go(
cls,
*,
model,
mode: str,
pred_flip_mode: bool,
require_uncond: bool,
start_sigma: float,
end_sigma: float,
blend_mode: str,
blend_strength: float,
blend_scale_mode: str,
blend_scale_offset: float,
blend_scale_min: float,
blend_scale_max: float,
operation_1=None,
operation_2=None,
operation_3=None,
operation_4=None,
operation_5=None,
) -> tuple:
if mode == "model_input":
if require_uncond:
raise ValueError(
"require_uncond does not make sense for the model_input mode.",
)
if pred_flip_mode:
raise ValueError(
"pred_flip does not make sense for the model_input mode.",
)
model = model.clone()
operations = tuple(
o
for o in (operation_1, operation_2, operation_3, operation_4, operation_5)
if o is not None
)
if not operations:
return (model,)
ms = model.get_model_object("model_sampling")
post_cfg_mode = mode in {"denoised", "denoised_sub_uncond"}
blend_function = utils.BLENDING_MODES[blend_mode]
sigma_max, sigma_min = (
ms.sigma_max.detach().item(),
ms.sigma_min.detach().item(),
)
if start_sigma < 0:
start_sigma = sigma_max
start_sigma = max(sigma_min, min(sigma_max, start_sigma))
end_sigma = max(sigma_min, min(sigma_max, end_sigma))
if end_sigma > start_sigma:
start_sigma, end_sigma = end_sigma, start_sigma
if start_sigma == end_sigma:
blend_scale_mode = "none"
orig_mode = mode
def patch(args: dict) -> torch.Tensor:
nonlocal mode
x = args["input"]
sigma_t = args["sigma"]
sigma_t_max = sigma_t.max()
if sigma_t.numel() > 1:
shape_pad = (1,) * (x.ndim - sigma_t.ndim)
sigma_t = sigma_t.reshape(sigma_t.shape[0], *shape_pad)
sigma = sigma_t_max.detach().item()
enabled = end_sigma <= sigma <= start_sigma
conds_out = args.get("conds_out", ())
uncond = (
args.get("uncond_denoised")
if post_cfg_mode
else (conds_out[1] if len(conds_out) > 1 else None)
)
if uncond is None and (
require_uncond
or mode in {"uncond", "uncond_sub_cond", "denoised_sub_uncond"}
):
enabled = False
if not enabled:
if mode == "model_input":
return x
return args["denoised"] if post_cfg_mode else conds_out
cond = conds_out[0] if not post_cfg_mode and len(conds_out) else None
if uncond is None and mode.endswith("_sub_uncond"):
mode = orig_mode.split("_", 1)[0]
else:
mode = orig_mode
if mode == "model_input":
t1 = x
t2 = None
elif mode in {"cond", "cond_sub_uncond"}:
t1 = cond
t2 = uncond if mode == "cond_sub_uncond" else None
elif mode in {"uncond", "uncond_sub_cond"}:
t1 = uncond
t2 = cond if mode == "uncond_sub_cond" else None
else:
t1 = args["denoised"]
t2 = uncond if mode == "denoised_sub_uncond" else None
t1_orig = t1
if pred_flip_mode:
t1 = (x - t1) / sigma_t
if t2 is not None:
t2 = (x - t2) / sigma_t
curr_blend = blend_strength * cls.get_blend_scaling(
scale_mode=blend_scale_mode,
offset=blend_scale_offset,
min_pct=blend_scale_min,
max_pct=blend_scale_max,
model_sampling=args["model"].model_sampling,
start_sigma=start_sigma,
end_sigma=end_sigma,
sigma=max(sigma_min, min(sigma, sigma_max)),
sigma_t_max=sigma_t_max.clamp(sigma_min, sigma_max),
)
result = t1 - t2 if t2 is not None else t1.clone()
for operation in operations:
result = operation(result)
if t2 is not None:
result += t2
if pred_flip_mode:
result = x - sigma_t * result
if curr_blend != 1:
result = blend_function(t1_orig, result, curr_blend)
if post_cfg_mode or mode == "model_input":
return result
conds_out = conds_out.copy()
conds_out[0 if mode.startswith("cond") else 1] = result
return conds_out
if post_cfg_mode:
model.set_model_sampler_post_cfg_function(patch)
elif mode == "model_input":
def patch_wrapper(apply_model, args: dict) -> torch.Tensor:
timestep = args["timestep"]
patch_args = args | {"sigma": timestep, "model": model.model}
return apply_model(patch(patch_args), timestep, **args["c"])
model.set_model_unet_function_wrapper(patch_wrapper)
else:
model.set_model_sampler_pre_cfg_function(patch)
return (model,)
class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
DESCRIPTION = "Allows applying a quantile normalization function to the latent during sampling. Can be used with Sonar SonarApplyLatentOperationCFG. The just copies most of the parameters from the other quantile normalization node where it talks to 'noise', this will apply to whatever you're applying the latent operation to (denoised, uncond, etc)."
RETURN_TYPES = ("LATENT_OPERATION",)
CATEGORY = "latent/advanced/operations"
@classmethod
def INPUT_TYPES(cls):
result = super().INPUT_TYPES()
result.pop("optional", None)
reqparams = result["required"]
for k in ("custom_noise", "normalize", "normalize_noise", "factor"):
reqparams.pop(k, None)
return result
@classmethod
def go(
cls,
*,
quantile: float,
dim: str,
flatten: bool,
norm_power: float,
norm_factor: float,
strategy: str,
):
def operation(latent: torch.Tensor, **_kwargs: dict) -> torch.Tensor:
return utils.quantile_normalize(
latent,
quantile=quantile,
dim=None if dim == "global" else int(dim),
flatten=flatten,
nq_fac=norm_factor,
pow_fac=norm_power,
strategy=strategy,
)
return (operation,)
class CustomNOISE:
def __init__(
self,
@@ -3102,29 +3619,32 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode):
)
else:
sampler_kwargs |= extra_params
sampler_function = functools.update_wrapper(
functools.partial(
self.sampler_function,
override_sampler_cfg={
"sampler": sampler,
"noise_type": NoiseType[noise_type.upper()]
if noise_type not in {None, "DEFAULT"}
else None,
"custom_noise": custom_noise_opt,
"sampler_kwargs": sampler_kwargs,
"cpu_noise": cpu_noise,
"normalize": normalize,
},
),
sampler.sampler_function,
)
return (
samplers.KSAMPLER(
self.sampler_function,
extra_options=sampler.extra_options
| {
"override_sampler_cfg": {
"sampler": sampler,
"noise_type": NoiseType[noise_type.upper()]
if noise_type not in {None, "DEFAULT"}
else None,
"custom_noise": custom_noise_opt,
"sampler_kwargs": sampler_kwargs,
"cpu_noise": cpu_noise,
"normalize": normalize,
},
},
inpaint_options=sampler.inpaint_options | {},
sampler_function,
extra_options=sampler.extra_options.copy(),
inpaint_options=sampler.inpaint_options.copy(),
),
)
@classmethod
@staticmethod
def sampler_function(
cls,
model,
x,
sigmas,
@@ -3133,7 +3653,7 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode):
noise_sampler: Callable | None = None,
extra_args: dict[str, Any] | None = None,
**kwargs: dict[str, Any],
):
) -> torch.Tensor:
if not override_sampler_cfg:
raise ValueError("Override sampler config missing!")
if extra_args is None:
@@ -3214,8 +3734,11 @@ NODE_CLASS_MAPPINGS = {
"SonarResizedNoise": SonarResizedNoiseNode,
"SonarWaveletNoise": SonarWaveletNoiseNode,
"SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode,
"SonarRippleFilteredNoise": SonarRippleFilteredNoiseNode,
"SonarQuantileFilteredNoise": SonarQuantileFilteredNoiseNode,
"SONAR_CUSTOM_NOISE to NOISE": SonarToComfyNOISENode,
"SonarApplyLatentOperationCFG": SonarApplyLatentOperationCFG,
"SonarLatentOperationQuantileFilter": SonarLatentOperationQuantileFilter,
}
+94
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import abc
import math
from functools import partial
from typing import Callable
@@ -1106,6 +1107,99 @@ class ChannelNoise(CustomNoiseItemBase):
return noise_sampler
class RippleFilteredNoise(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
noise,
mode: str,
dim: int,
flatten: bool,
offset: float,
amplitude_high: float,
amplitude_low: float,
period: float,
roll: float,
normalize_noise: float,
normalize,
):
super().__init__(
factor,
noise=noise.clone(),
mode=mode,
dim=dim,
flatten=flatten,
offset=offset,
amplitude_high=amplitude_high,
amplitude_low=amplitude_low,
period=period,
roll=roll,
normalize_noise=normalize_noise,
normalize=normalize,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
factor = self.factor
dim = self.dim
if dim < 0:
dim = x.ndim + dim
if dim < 0 or dim >= x.ndim:
raise ValueError("Dimension out of range")
dim_els = math.prod(x.shape[dim:]) if self.flatten else x.shape[dim]
mode_fun = torch.sin if self.mode.startswith("sin") else torch.cos
follow_sign = self.mode.endswith("_copysign")
scaler_shape = [1] * x.ndim
if self.flatten:
scaler_shape[dim:] = x.shape[dim:]
else:
scaler_shape[dim] = x.shape[dim]
scaler = mode_fun(
torch.linspace(
self.offset,
self.offset + math.pi * self.period,
steps=dim_els,
dtype=x.dtype,
device=x.device,
),
)
scaler = (
1.0
+ torch.where(
scaler < 0,
scaler * self.amplitude_low,
scaler * self.amplitude_high,
)
).reshape(scaler_shape)
ns = self.noise.make_noise_sampler(
x,
*args,
normalized=self.normalize_noise,
**kwargs,
)
roll = self.roll
normalize = self.get_normalize("normalize", normalized)
counter = 0
def noise_sampler(s, sn):
nonlocal counter
noise = ns(s, sn)
to_roll = int(roll * counter)
counter += 1
scaler_curr = scaler.roll(to_roll, dims=dim)
result = scale_noise(noise, factor, normalized=normalize).mul_(
scaler_curr,
)
return result.copysign(1.0 - scaler_curr) if follow_sign else result
return noise_sampler
class BlendedNoise(CustomNoiseItemBase):
def __init__(
self,
+194 -45
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import math
from functools import partial
import torch
from comfy.model_management import device_supports_non_blocking
@@ -86,30 +87,150 @@ def tensor_to(
return tensor.to(dest, non_blocking=non_blocking)
def _quantile_norm_scaledown(noise: torch.Tensor, nq: torch.Tensor) -> torch.Tensor:
mv = noise.abs().max().detach().item()
return noise if mv == 0 else torch.where(noise.abs() > nq, noise * (nq / mv), noise)
def _quantile_norm_scaledown(
noise: torch.Tensor,
nq: torch.Tensor,
*,
dim,
**_kwargs: dict,
) -> torch.Tensor:
noiseabs = noise.abs()
mv = noiseabs.max(dim=dim, keepdim=True).clamp(min=1e-06)
return noise if mv == 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,
)
quantile_handlers = {
"clamp": lambda noise, nq: noise.clamp(-nq, nq),
"clamp": lambda noise, nq, **_kwargs: noise.clamp(-nq, nq),
"scale_down": _quantile_norm_scaledown,
"tanh": lambda noise, nq: noise.tanh().mul_(nq.abs()),
"tanh_outliers": lambda noise, nq: torch.where(
"tanh": lambda noise, nq, **_kwargs: noise.tanh().mul_(nq.abs()),
"tanh_outliers": lambda noise, nq, **_kwargs: torch.where(
noise.abs() > nq,
noise.tanh().mul_(nq.abs()),
noise,
),
"sigmoid": lambda noise, nq: noise.sigmoid().mul_(nq.abs()).copysign(noise),
"sigmoid_outliers": lambda noise, nq: torch.where(
"sigmoid": lambda noise, nq, **_kwargs: noise.sigmoid()
.mul_(nq.abs())
.copysign(noise),
"sigmoid_outliers": lambda noise, nq, **_kwargs: torch.where(
noise.abs() > nq,
noise.sigmoid().mul_(nq.abs()).copysign(noise),
noise,
),
"tenth": lambda noise, nq: torch.where(noise.abs() > nq, noise * 0.1, noise),
"half": lambda noise, nq: torch.where(noise.abs() > nq, noise * 0.5, noise),
"zero": lambda noise, nq: torch.where(noise.abs() > nq, 0, noise),
"reverse_zero": lambda noise, nq: torch.where(noise.abs() >= nq, noise, 0),
"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,
noise,
),
"half": lambda noise, nq, **_kwargs: torch.where(
noise.abs() > nq,
noise * 0.5,
noise,
),
"zero": lambda noise, nq, **_kwargs: torch.where(noise.abs() > nq, 0, noise),
"reverse_zero": lambda noise, nq, **_kwargs: torch.where(
noise.abs() >= nq,
noise,
0,
),
"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),
}
@@ -124,41 +245,66 @@ def quantile_normalize(
pow_fac: float = 0.5,
strategy: str = "clamp",
strategy_handler=None,
use_abs: bool = True,
) -> torch.Tensor:
if noise.numel() == 0:
return noise
if isinstance(quantile, (tuple, list)):
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,
use_abs=use_abs,
)
return noise
if quantile is None or quantile <= 0 or quantile >= 1:
return noise
orig_shape = noise.shape
if isinstance(quantile, (tuple, list)):
quantile = torch.tensor(
quantile,
device=noise.device,
dtype=noise.dtype,
if not use_abs:
pos_mask = noise >= 0
neg_mask = ~pos_mask
result = torch.zeros_like(noise)
result[pos_mask] = quantile_normalize(
noise=noise[pos_mask],
quantile=quantile,
dim=dim,
flatten=flatten,
nq_fac=nq_fac,
pow_fac=pow_fac,
strategy=strategy,
strategy_handler=strategy_handler,
use_abs=True,
)
qdim = dim
result[neg_mask] = quantile_normalize(
noise=noise[neg_mask],
quantile=quantile,
dim=dim,
flatten=flatten,
nq_fac=nq_fac,
pow_fac=pow_fac,
strategy=strategy,
strategy_handler=strategy_handler,
use_abs=True,
)
return result
orig_shape = noise.shape
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 qdim in {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",
)
flatnoise = noise.flatten(start_dim=dim)
else:
flatdim = None
flatten = False
flatnoise = noise
nq = torch.quantile(
(noise if flatdim is None else noise.flatten(start_dim=flatdim)).abs(),
flatnoise.abs(),
quantile,
dim=-1,
dim=-1 if flatten else dim,
keepdim=True,
)
nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim)
nq = nq.mul_(nq_fac).reshape(*nq_shape)
nq = nq.mul_(nq_fac)
handler = (
quantile_handlers.get(strategy)
if strategy_handler is None
@@ -166,13 +312,16 @@ def quantile_normalize(
)
if handler is None:
raise ValueError("Unknown strategy")
noise = handler(noise, nq)
noise = noise.abs().pow(pow_fac).copysign(noise)
if flatdim is not None and qdim in {2, 3}:
return (
noise.reshape(tempshape).movedim(1, qdim).reshape(orig_shape).contiguous()
)
return noise
noise = handler(
flatnoise,
nq,
orig_noise=noise,
dim=dim,
flatten=flatten,
)
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 normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)):