diff --git a/changelog.md b/changelog.md index ee5732c..ae8a8b3 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,12 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 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). +* Added actual wavelet noise based on https://en.wikipedia.org/wiki/Wavelet_noise . +* Added `reverse_zero`, `scale_down`, `tanh`, `tanh_outliers`, `sigmoid` and `sigmoid_outliers` quantile normalization limit modes. + ## 20250602 * Fixed broken calculation for Collatz noise. diff --git a/py/nodes.py b/py/nodes.py index 989a05c..46b2283 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -308,7 +308,7 @@ class SonarNoiseImageNode(metaclass=IntegratedNode): "noise_multiplier": ( "FLOAT", { - "default": 1.0, + "default": 0.5, "step": 0.001, "min": -1000.0, "max": 1000.0, @@ -1641,36 +1641,49 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): "adjust_scale": ( "BOOLEAN", { - "default": True, - }, - ), - "use_initial": ( - "BOOLEAN", - { - "default": True, - }, - ), - "iteration_sign_flipping": ( - "BOOLEAN", - { - "default": True, + "default": False, + "tooltip": "When enabled, the output will be normalized to values between -1 and 1 using the last two dimensions (if there are four or more), otherwise dimensions after the first.", }, ), "chain_length": ( "STRING", { "default": "1, 1, 2, 2, 3, 3", - "tooltip": "Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap.", + "tooltip": "Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap. Controls the length of Collatz chains. Note: Using a high chain length may be very slow, especially if combined with many iterations.", + }, + ), + "chain_offset": ( + "INT", + { + "default": 5, + "min": 0, + "max": 10000, + "tooltip": "Uses values starting at the specified offset. Note: This entails generating chains of length chain_length + chain_offset, which may be quite slow if you use high values.", + }, + ), + "iterations": ( + "INT", + { + "default": 10, + "min": 1, + "max": 10000, + "tooltip": "Number of iterations to run. Warning: Collatz noise (my implementation, anyway) is EXTREMELY slow.", + }, + ), + "iteration_sign_flipping": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether we cycle between flipping the sign on the output from each iteration. May average out weirdness... Or make stuff weirder.", }, ), - "iterations": ("INT", {"default": 500, "min": 1, "max": 10000}), "rmin": ( "FLOAT", { "default": -8000.0, "min": -100000.0, "max": 100000.0, - "tooltip": "Going as low as -9500 should be safe.", + "tooltip": "Minimum value a chain can start with. Going as low as -9500 should be safe with float32.", }, ), "rmax": ( @@ -1679,10 +1692,9 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): "default": 8000.0, "min": -100000.0, "max": 100000.0, - "tooltip": "I don't recommend going over 9500 here as that is where the Collatz chain starts to reach values that can't be accurately represented with a 32bit float.", + "tooltip": "Maximum value a chain can start with. I don't recommend going over 9500 if you are using the float32 dtype here as that is where the Collatz chain starts to reach values that can't be accurately represented.", }, ), - "flatten": ("BOOLEAN", {"default": False}), "dims": ( "STRING", { @@ -1690,13 +1702,128 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): "tooltip": "Comma-separated list of dimensions. Cannot be empty. May be negative to count from the end of the list. Iterations will cycle through the list and wrap.", }, ), - "variant": ( - "INT", + "flatten": ( + "BOOLEAN", { - "default": 2, - "min": 1, - "max": 2, - "tooltip": "Variant 1 may act like the original version, not the correct algorithm for Collatz though. Variant 2 is (hopefully) more correct. There will likely be more variants in the future.", + "default": False, + "tooltip": "Controls whether dimensions past the current one selected from the dims parameter will get flattened.", + }, + ), + "output_mode": ( + ( + "values", + "ratios", + "mults", + "adds", + "seed_x_mults", + "seed_x_adds", + "noise_x_ratios", + "noise_x_mults", + "noise_x_adds", + ), + { + "default": "values", + }, + ), + "quantile": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "tooltip": "The initial output of each iteration will be run through quantile normalization. Setting the parameter to 0 or 1 will disable quantile normalization.", + }, + ), + "quantile_strategy": ( + tuple(utils.quantile_handlers.keys()), + { + "default": "clamp", + "tooltip": "Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", + }, + ), + "noise_dtype": ( + ("float32", "float64", "float16", "bfloat16"), + { + "default": "float32", + "tooltip": "Generally should be left at the default. Only float32 and float64 will work if you have quantile normalization enabled.", + }, + ), + "even_multiplier": ( + "FLOAT", + { + "default": 0.5, + "min": -10000.0, + "max": 1000.0, + "tooltip": "Multiplier to use when the previous link in the chain is even. Collatz uses 0.5 (divides by two) here.", + }, + ), + "even_addition": ( + "FLOAT", + { + "default": 0.0, + "min": -10000.0, + "max": 1000.0, + "tooltip": "Value to add when the previous link in the chain is even. Collatz uses 0 here.", + }, + ), + "odd_multiplier": ( + "FLOAT", + { + "default": 3.0, + "min": -10000.0, + "max": 1000.0, + "tooltip": "Multiplier to use when the previous link in the chain is odd. Collatz uses 3 here.", + }, + ), + "odd_addition": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 1000.0, + "tooltip": "Value to add when the previous link in the chain is odd. Collatz uses 1 here.", + }, + ), + "integer_math": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether the results during chain generation get truncated to an integer value or not. Should be enabled if you actually want to generate accurate Collatz chains.", + }, + ), + "add_preserves_sign": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether additions use the same sign as the item they're being added to.", + }, + ), + "break_loops": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Controls whether the chain resets back to the seed value once it reaches 1 or 0. Generally should be left enabled, otherwise the chain will oscillate between only a few values for the rest of the length (at least with the Collatz rules).", + }, + ), + "seed_mode": ( + ("default", "force_odd", "force_even"), + { + "default": "default", + "tooltip": "Default mode just uses whatever the original seed value was. force_odd/force_even will force it to the specified parity by adding one if it doesn't match. Starting from odd seeds might result in longer chains. Enabling the force modes may cause the initial seeds to exceed rmax by one.", + }, + ), + } + result["optional"] |= { + "seed_custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": f"Optional custom noise to use for initial values for Collatz chains. May be slow as it will generate noise according to the original input size and then crop it. Does this noise type have enough warnings about it being slow? Yeah. Connecting something here will probably make it even slower!\n{NOISE_INPUT_TYPES_HINT}", + }, + ), + "mix_custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": f"Optional custom noise to use with the output modes starting with 'noise'.\n{NOISE_INPUT_TYPES_HINT}", }, ), } @@ -1712,7 +1839,6 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): factor: float, rescale: float, adjust_scale: bool, - use_initial: bool, iteration_sign_flipping: bool, chain_length: int, iterations: int, @@ -1720,7 +1846,21 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): rmax: float, flatten: bool, dims: str, - variant: int, + output_mode: str, + noise_dtype: str, + quantile: float, + quantile_strategy: str, + integer_math: bool, + add_preserves_sign: bool, + even_multiplier: float, + even_addition: float, + odd_multiplier: float, + odd_addition: float, + chain_offset: int, + seed_mode: str, + break_loops: bool, + seed_custom_noise: object | None = None, + mix_custom_noise: object | None = None, sonar_custom_noise_opt=None, ): if rmin > rmax: @@ -1731,7 +1871,6 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): rescale=rescale, sonar_custom_noise_opt=sonar_custom_noise_opt, adjust_scale=adjust_scale, - use_initial=use_initial, iteration_sign_flipping=iteration_sign_flipping, chain_length=tuple(int(i) for i in chain_length.split(",")), iterations=iterations, @@ -1739,7 +1878,26 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): rmax=rmax, flatten=flatten, dims=dims, - variant=variant, + output_mode=output_mode, + quantile=quantile, + quantile_strategy=quantile_strategy, + integer_math=integer_math, + add_preserves_sign=add_preserves_sign, + even_multiplier=even_multiplier, + even_addition=even_addition, + odd_multiplier=odd_multiplier, + odd_addition=odd_addition, + chain_offset=chain_offset, + break_loops=break_loops, + seed_mode=seed_mode, + noise_dtype={ + "float32": torch.float32, + "float64": torch.float64, + "float16": torch.float16, + "bfloat16": torch.bfloat16, + }.get(noise_dtype, torch.float32), + seed_custom_noise=seed_custom_noise, + mix_custom_noise=mix_custom_noise, ) @@ -1816,10 +1974,10 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase): }, ), "strategy": ( - ("clamp", "half", "tenth", "zero"), + tuple(utils.quantile_handlers.keys()), { "default": "clamp", - "tooltip": "Determines how to treat outliers.", + "tooltip": "Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", }, ), } @@ -2158,6 +2316,192 @@ class SonarWaveletFilteredNoiseNode( ) +class SonarWaveletNoiseNode( + SonarCustomNoiseNodeBase, + SonarNormalizeNoiseNodeMixin, +): + DESCRIPTION = "Custom noise type that allows generating wavelet noise. Very simple explanation of how a single octave works:\n1) Generate some noise.\n2) Scale it down 50%.\n3) Scale it back up to the original size.\n4) Subtract the scaled noise from the original noise.\nScaling the noise down and then back up blurs it, so this is essentially sharpening the noise." + + @classmethod + def INPUT_TYPES(cls): + result = super().INPUT_TYPES() + result["required"] |= { + "octaves": ( + "INT", + { + "default": 4, + "min": -100, + "max": 100, + "tooltip": "Number of octaves to generate. You can use a negative number here to run the octaves in reverse order though it may produce weird results/not work very well.", + }, + ), + "octave_height_factor": ( + "FLOAT", + { + "default": 0.5, + "min": 0.001, + "max": 10000.0, + "tooltip": "Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", + }, + ), + "octave_width_factor": ( + "FLOAT", + { + "default": 0.5, + "min": 0.001, + "max": 10000.0, + "tooltip": "Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", + }, + ), + "octave_scale_mode": ( + utils.UPSCALE_METHODS, + { + "tooltip": "Scaling mode used within each octave to produce the scaled noise. By default this will be scaling down that octave's noise.", + "default": "adaptive_avg_pool2d", + }, + ), + "octave_rescale_mode": ( + utils.UPSCALE_METHODS, + { + "tooltip": "Scaling mode used within each octave to scale the noise back up to that octave's original size.", + "default": "bilinear", + }, + ), + "post_octave_rescale_mode": ( + utils.UPSCALE_METHODS, + { + "tooltip": "Scaling mode used to scale the output of an octave back up to the actual latent size.", + "default": "bilinear", + }, + ), + "initial_amplitude": ( + "FLOAT", + { + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Basically the strength an octave gets added to the total. This will be scaled by persistance after each octave.", + }, + ), + "persistence": ( + "FLOAT", + { + "default": 0.5, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Multiplier applied to amplitude after each octave. 0.5 means the first octave uses initial_amplitude, the second uses half of that and so on.", + }, + ), + "height_factor": ( + "FLOAT", + { + "default": 2.0, + "min": 0.001, + "max": 10000.0, + "tooltip": "Scaling factor for height, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", + }, + ), + "width_factor": ( + "FLOAT", + { + "tooltip": "Scaling factor for width, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", + "default": 2.0, + "min": 0.001, + "max": 10000.0, + }, + ), + "update_blend": ( + "FLOAT", + { + "tooltip": "Controls how original_noise - scaled_noise is blended with original_noise. The default is to use 100% original_noise - scaled_noise.", + "default": 1.0, + "min": -10000.0, + "max": 10000.0, + }, + ), + "update_blend_mode": ( + ("simple_add", *utils.BLENDING_MODES.keys()), + { + "default": "lerp", + "tooltip": "Controls how the enhanced noise from each octave is blended with that octave's raw noise. With normal wavelet noise there's no blending and you use 100% enhanced noise.", + }, + ), + "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.", + }, + ), + } + result["optional"] |= { + "custom_noise": ( + WILDCARD_NOISE, + { + "tooltip": f"Optional: Custom noise input. If unconnected will default to Gaussian noise. Note: When connected, the noise for all octaves will be generated at the maximum scale and then cropped which may be slow.\n{NOISE_INPUT_TYPES_HINT}", + }, + ), + } + return result + + @classmethod + def get_item_class(cls): + return noise.AdvancedWaveletNoise + + def go( + self, + *, + factor, + rescale, + normalize, + octaves: int, + octave_height_factor: float, + octave_width_factor: float, + octave_scale_mode: str, + octave_rescale_mode: str, + post_octave_rescale_mode: str, + initial_amplitude: float, + persistence: float, + height_factor: float, + width_factor: float, + update_blend: float, + update_blend_mode: str, + normalize_noise: bool, + custom_noise=None, + sonar_custom_noise_opt=None, + ): + if persistence == 0 or initial_amplitude == 0 or octaves == 0: + raise ValueError( + "Persistence, initial amplitude and octaves must be non-zero", + ) + return super().go( + factor, + rescale=rescale, + sonar_custom_noise_opt=sonar_custom_noise_opt, + octaves=octaves, + octave_height_factor=octave_height_factor, + octave_width_factor=octave_width_factor, + octave_scale_mode=octave_scale_mode, + octave_rescale_mode=octave_rescale_mode, + post_octave_rescale_mode=post_octave_rescale_mode, + initial_amplitude=initial_amplitude, + persistence=persistence, + height_factor=height_factor, + width_factor=width_factor, + update_blend=update_blend, + update_blend_function=utils.BLENDING_MODES[update_blend_mode], + normalize=self.get_normalize(normalize), + normalize_noise=normalize_noise, + custom_noise=custom_noise, + ) + + class CustomNOISE: def __init__( self, @@ -2868,6 +3212,7 @@ NODE_CLASS_MAPPINGS = { "SonarChannelNoise": SonarChannelNoiseNode, "SonarBlendedNoise": SonarBlendedNoiseNode, "SonarResizedNoise": SonarResizedNoiseNode, + "SonarWaveletNoise": SonarWaveletNoiseNode, "SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode, "SonarQuantileFilteredNoise": SonarQuantileFilteredNoiseNode, "SONAR_CUSTOM_NOISE to NOISE": SonarToComfyNOISENode, diff --git a/py/noise.py b/py/noise.py index ee25ed5..25f4e4c 100644 --- a/py/noise.py +++ b/py/noise.py @@ -197,9 +197,10 @@ class NoiseSampler: seed: int | None = None, cpu: bool = False, transform: Callable = lambda t: t, - make_noise_sampler: Callable | None = None, normalized=False, factor: float = 1.0, + *, + make_noise_sampler: Callable, **kwargs, ): self.factor = factor @@ -322,7 +323,6 @@ class AdvancedDistroNoise(AdvancedNoiseBase): class AdvancedCollatzNoise(AdvancedNoiseBase): ns_factory_arg_keys = ( "adjust_scale", - "use_initial", "iteration_sign_flipping", "chain_length", "iterations", @@ -330,13 +330,109 @@ class AdvancedCollatzNoise(AdvancedNoiseBase): "rmax", "flatten", "dims", - "variant", + "output_mode", + "noise_dtype", + "quantile", + "quantile_strategy", + "integer_math", + "add_preserves_sign", + "even_multiplier", + "even_addition", + "odd_multiplier", + "odd_addition", + "chain_offset", + "seed_mode", + "break_loops", ) @property def ns_factory(self): return CollatzNoiseGenerator + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + seed_ns = ( + self.seed_custom_noise.make_noise_sampler( + x, + *args, + normalized=False, + **kwargs, + ) + if self.seed_custom_noise is not None + else None + ) + mix_ns = ( + self.mix_custom_noise.make_noise_sampler( + x, + *args, + normalized=False, + **kwargs, + ) + if self.mix_custom_noise is not None + and self.output_mode.startswith("noise_") + else None + ) + return super().make_noise_sampler( + x, + *args, + normalized=normalized, + seed_noise_sampler=seed_ns, + mix_noise_sampler=mix_ns, + ) + + +class AdvancedWaveletNoise(AdvancedNoiseBase): + ns_factory_arg_keys = ( + "octave_scale_mode", + "octave_rescale_mode", + "post_octave_rescale_mode", + "initial_amplitude", + "persistence", + "octaves", + "octave_height_factor", + "octave_width_factor", + "height_factor", + "width_factor", + "min_height", + "min_width", + "update_blend", + "update_blend_function", + ) + + @property + def ns_factory(self): + return WaveletNoiseGenerator + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + if x.ndim < 4: + raise ValueError("Can only handle 4+ dimensional latents") + height, width = x.shape[-2:] + result = super().make_noise_sampler(x, *args, normalized=normalized, **kwargs) + wavelet_ng = result.noise_sampler + max_height = ( + int(max(height, *(od.height for od in wavelet_ng.octave_data))) + if wavelet_ng.octave_data + else height + ) + max_width = ( + int(max(width, *(od.width for od in wavelet_ng.octave_data))) + if wavelet_ng.octave_data + else width + ) + internal_ns = ( + self.custom_noise.make_noise_sampler( + x.new_zeros(*x.shape[:-2], max_height, max_width) + if max_width != width or max_height != height + else x, + *args, + normalized=self.normalize_noise, + **kwargs, + ) + if self.custom_noise is not None + else None + ) + wavelet_ng.set_internal_noise_sampler(internal_ns) + return result + class CompositeNoise(CustomNoiseItemBase): def __init__( @@ -1240,7 +1336,7 @@ class WaveletFilteredNoise(CustomNoiseItemBase): ns_kwargs = getattr(self, "ns_kwargs", {}).copy() # print("WF:NS KWARGS", ns_kwargs) kwargs |= ns_kwargs - ns = WaveletNoiseGenerator( + ns = WaveletFilterNoiseGenerator( x, *args, sigma_min=sigma_min, @@ -1528,9 +1624,10 @@ class PatternBreakNoise(CustomNoiseItemBase): *, noise, detail_level: float, - blend_mode: str, percentage: float, restore_scale: bool, + blend_mode: str = "lerp", + blend_function=None, ): super().__init__( factor, @@ -1538,7 +1635,7 @@ class PatternBreakNoise(CustomNoiseItemBase): detail_level=detail_level, percentage=percentage, restore_scale=restore_scale, - blend_function=utils.BLENDING_MODES[blend_mode], + blend_function=blend_function or utils.BLENDING_MODES[blend_mode], ) def clone_key(self, k): @@ -1663,8 +1760,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = { MixedNoiseGenerator, name="onef_pinkish_mix", noise_mix=( - (OneFNoiseGenerator, {"alpha": 0.5}, lambda t: t.mul_(-1.0)), - (OneFNoiseGenerator, {"alpha": 0.5}, None), + (OneFNoiseGenerator, {"alpha": -0.5}, lambda t: t.mul_(-1.0)), + (OneFNoiseGenerator, {"alpha": -0.5}, None), ), output_fun=lambda t: t.mul_(0.5), ), diff --git a/py/noise_generation.py b/py/noise_generation.py index a983506..4d32a42 100644 --- a/py/noise_generation.py +++ b/py/noise_generation.py @@ -3,8 +3,7 @@ from __future__ import annotations import math from enum import Enum, auto -from functools import partial -from typing import Callable +from typing import TYPE_CHECKING, Callable, ClassVar, NamedTuple import torch from comfy.k_diffusion import sampling @@ -18,6 +17,8 @@ try: except ImportError: HAVE_WAVELETS = False +from comfy.model_management import throw_exception_if_processing_interrupted + from . import utils from .utils import ( fallback, @@ -27,6 +28,9 @@ from .utils import ( tensor_to, ) +if TYPE_CHECKING: + from collections.abc import Sequence + # ruff: noqa: D413, D417, D212, ANN002, ANN003 @@ -106,7 +110,6 @@ class NoiseGenerator: setattr(self, k, kwarg_params.pop(k)) self.options = kwarg_params self.update_x(x) - # print("CREATE NG", self, kwargs) @classmethod def ng_params(cls): @@ -131,14 +134,25 @@ class NoiseGenerator: self.layout = x.layout self.dtype = x.dtype - def rand_like(self, *, fun=torch.randn, cpu=None, to_device=True): - cpu = cpu if cpu is not None else self.cpu + def rand_like( + self, + *, + fun=torch.randn, + cpu=None, + to_device=True, + shape=None, + dtype=None, + layout=None, + device=None, + generator=None, + ): + cpu = fallback(cpu, self.cpu) noise = fun( - *self.shape, - generator=self.generator, - dtype=self.dtype, - layout=self.layout, - device=self.gen_device, + *fallback(shape, self.shape), + generator=fallback(generator, self.generator), + dtype=fallback(dtype, self.dtype), + layout=fallback(layout, self.layout), + device=fallback(device, "cpu" if cpu else self.gen_device), ) if to_device and noise.device != self.device: noise = tensor_to(noise, self.device) @@ -189,8 +203,10 @@ class FramesToChannelsNoiseGenerator(NoiseGenerator): self.width, ) - def rand_like(self, *args, **kwargs): - noise = super().rand_like(*args, **kwargs) + def rand_like(self, *args, shape=None, **kwargs): + noise = super().rand_like(*args, shape=shape, **kwargs) + if shape is not None: + return noise adjusted_shape = self.get_adjusted_shape() if noise.shape != adjusted_shape: return noise.reshape(*adjusted_shape) @@ -1280,8 +1296,8 @@ class PowerOldNoiseGenerator(NoiseGenerator): # Idea from https://github.com/ClownsharkBatwing/RES4LYF/ (wave and mode defaults also from that source) -class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator): - name = "wavelet" +class WaveletFilterNoiseGenerator(FramesToChannelsNoiseGenerator): + name = "waveletfilter" MIN_DIMS = 4 MAX_DIMS = 5 @@ -1370,101 +1386,396 @@ class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator): return result[tuple(slice(0, dl) for dl in noise.shape)] -class CollatzNoiseGenerator(NoiseGenerator): - name = "collatz" +class WaveletNoiseOctave(NamedTuple): + octave: int + height: int + width: int + amplitude: float + total_amplitude: float + + +class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator): + name = "wavelet" + MIN_DIMS = 4 + MAX_DIMS = 5 @classmethod def ng_params(cls): return super().ng_params() | { - "adjust_scale": True, - "use_initial": True, + "octave_scale_mode": "adaptive_avg_pool2d", + "octave_rescale_mode": "bilinear", + "post_octave_rescale_mode": "bilinear", + "initial_amplitude": 1.0, + "persistence": 0.5, + "octaves": 4, + "octave_height_factor": 0.5, + "octave_width_factor": 0.5, + "height_factor": 2.0, + "width_factor": 2.0, + "min_height": 4, + "min_width": 4, + "update_blend": 1.0, + "update_blend_function": torch.lerp, + "noise_sampler": None, + } + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.set_octave_data() + + def set_internal_noise_sampler(self, noise_sampler: object) -> None: + self.noise_sampler = noise_sampler + + def set_octave_data(self) -> tuple: + adjusted_shape = self.get_adjusted_shape() + height, width = adjusted_shape[-2:] + amplitude = self.initial_amplitude + total_amplitude = 0.0 + curr_height, curr_width = height, width + octave_data = [] + is_reverse = self.octaves < 0 + octaves = ( + range(self.octaves) + if not is_reverse + else reversed(range(abs(self.octaves))) + ) + for octave in octaves: + curr_height /= self.height_factor**octave + curr_width /= self.width_factor**octave + if ( + amplitude == 0 + or curr_height < self.min_height + or curr_width < self.min_width + or curr_height * self.octave_height_factor < 1 + or curr_width * self.octave_width_factor < 1 + ): + if is_reverse and not octave_data: + curr_height, curr_width = height, width + continue + break + total_amplitude += abs(amplitude) + octave_data.append( + WaveletNoiseOctave( + octave=octave, + height=curr_height, + width=curr_width, + amplitude=amplitude, + total_amplitude=total_amplitude, + ), + ) + amplitude *= self.persistence + if not octave_data or not total_amplitude: + raise ValueError("Unworkable parameters for wavelet noise") + self.octave_data = tuple(octave_data) + + def _generate_octave(self, *args: list, shape: Sequence) -> torch.Tensor: + height, width = shape[-2:] + noise = ( + self.noise_sampler(*args)[..., :height, :width].reshape(shape) + if self.noise_sampler + else self.rand_like(shape=(*shape[:-2], height, width)) + ) + scaled_height = int(max(1, height * self.octave_height_factor)) + scaled_width = int(max(1, width * self.octave_width_factor)) + scaled_noise = utils.scale_samples( + utils.scale_samples( + noise, + scaled_width, + scaled_height, + mode=self.octave_scale_mode, + ), + width=width, + height=height, + mode=self.octave_rescale_mode, + ) + return self.update_blend_function( + noise, + noise - scaled_noise, + self.update_blend, + ) + + def generate(self, *args: list) -> torch.Tensor: + adjusted_shape = self.get_adjusted_shape() + height, width = adjusted_shape[-2:] + curr_shape = list(adjusted_shape) + result = torch.zeros( + adjusted_shape, + device=self.device, + dtype=self.dtype, + layout=self.layout, + ) + for od in self.octave_data: + curr_shape[-2:] = (int(od.height), int(od.width)) + octave_output = self._generate_octave(*args, shape=curr_shape) + if octave_output.shape != result.shape: + octave_output = utils.scale_samples( + octave_output, + width, + height, + mode=self.post_octave_rescale_mode, + ) + result += octave_output.mul_(od.amplitude) + if self.octave_data[-1].total_amplitude != 0: + result /= self.octave_data[-1].total_amplitude + return self.fix_output_frames(result) + + +class CollatzNoiseGenerator(NoiseGenerator): + name = "collatz" + + chain_cache: ClassVar[dict] = {} + + @classmethod + def ng_params(cls): + return super().ng_params() | { + "adjust_scale": False, "iteration_sign_flipping": True, "chain_length": (1, 1, 2, 2, 3, 3), - "iterations": 500, + "iterations": 10, "rmin": -8000.0, "rmax": 8000.0, "flatten": False, "dims": (-1, -1, -2, -2), - "variant": 2, + # values, ratios, mults, adds + # seed_x_ratios, seed_x_mults, seed_x_adds + # noise_x_ratios, noise_x_mults, noise_x_adds + "output_mode": "values", + "quantile": 0.5, + "quantile_strategy": "clamp", + "noise_dtype": torch.float32, + "integer_math": True, + "even_multiplier": 0.5, + "even_addition": 0.0, + "odd_multiplier": 3.0, + "odd_addition": 1.0, + "add_preserves_sign": True, + "chain_offset": 5, + "break_loops": True, + "seed_mode": "default", + "seed_noise_sampler": None, + "mix_noise_sampler": None, } @staticmethod - def _get_iter_slices(n_dims, dim, offset, stride) -> tuple: - return tuple( - slice(None) if didx != dim else slice(offset, None, stride) - for didx in range(n_dims) - ) + def _get_iter_slices(n_dims, dim, idx, stride) -> list: + result = [slice(None)] * n_dims + result[dim] = slice(idx, None, stride) + return result - def _generate_iteration( + def _generate_iteration( # noqa: PLR0914 self, - *, + *args, dim: int, chain_length: int, flatten: False, shape=None, - integer_division=True, ): dtype, device = self.dtype, self.device out_shape = shape = fallback(shape, self.shape) if dim >= len(shape): raise ValueError("Requested dimension out of range") rmin, rmax = self.rmin, self.rmax + emul, eadd = self.even_multiplier, self.even_addition + omul, oadd = self.odd_multiplier, self.odd_addition + keepsign = self.add_preserves_sign + intmode = self.integer_math rmaxsubmin = rmax - rmin if flatten: shape = torch.Size((*shape[:dim], math.prod(shape[dim:]))) size = shape[dim] chain_length = min(size, chain_length) n_chunks = math.ceil(size / chain_length) - result_shape = tuple( - (chain_length * n_chunks) if idx == dim else sz - for idx, sz in enumerate(shape) - ) - chunk_shape = tuple( - n_chunks if idx == dim else sz for idx, sz in enumerate(shape) - ) - result = torch.zeros(result_shape, dtype=torch.float32, device=device) - noise = ( - torch.rand( - chunk_shape, - generator=self.generator, - dtype=torch.float32, - device=self.gen_device, - layout=self.layout, + chain_length += self.chain_offset + result_shape = list(shape) + chunk_shape = result_shape.copy() + result_shape[dim] = chain_length * n_chunks + chunk_shape[dim] = n_chunks + result = torch.zeros(result_shape, dtype=self.noise_dtype, device=device) + adds, muls = result.clone(), result.clone() + if self.seed_noise_sampler is not None: + orig_noise = self.seed_noise_sampler(*args)[ + tuple(slice(None, sz) for sz in chunk_shape) + ].to(result) + if flatten: + orig_noise = orig_noise.flatten(start_dim=dim) + orig_noise = normalize_to_scale( + orig_noise[tuple(slice(None, sz) for sz in chunk_shape)], + 1e-06, + 1.0, + dim=tuple(range(1, len(chunk_shape))), + ) + else: + orig_noise = self.rand_like( + fun=torch.rand, + shape=chunk_shape, + dtype=result.dtype, + ) + noise = orig_noise * (rmaxsubmin + 1) + rmin + # Derp. + noise = torch.where(noise == 0, noise.max() / noise.numel(), noise) + if self.seed_mode != "default": + noise = torch.where( + (noise % 2.0) < 1 + if self.seed_mode == "force_odd" + else (noise % 2.0) >= 1, + noise + 1, + noise, ) - .mul_(rmaxsubmin + 1) - .add_(rmin) - ) - # noise = torch.where((noise % 2.0) < 1, noise + 1, noise) if noise.device != self.device: noise = tensor_to(noise, self.device) + slice_0 = self._get_iter_slices(result.ndim, dim, 0, chain_length) for chainidx in range(chain_length): - if chainidx == 0 and self.use_initial: - result[self._get_iter_slices(result.ndim, dim, 0, chain_length)] = noise + if chainidx == 0: + muls[slice_0] = 1.0 + result[slice_0] = noise continue - chunk = ( - noise - if chainidx == 0 - else result[ - self._get_iter_slices(result.ndim, dim, chainidx - 1, chain_length) - ] + slice_curr = self._get_iter_slices(result.ndim, dim, chainidx, chain_length) + slice_prev = self._get_iter_slices( + result.ndim, + dim, + chainidx - 1, + chain_length, ) - result[self._get_iter_slices(result.ndim, dim, chainidx, chain_length)] = ( + prev = result[slice_prev] + prev_trunc = utils.trunc_decimals(prev, 2) + need_reset = ( + ((prev_trunc >= 1.0) & (prev_trunc < 1.001)) + | (prev_trunc.abs() < 0.001) + if self.break_loops + else False + ) + prev_evens = prev % 2 < 1.0 + prev_adds, prev_muls = adds[slice_prev], muls[slice_prev] + muls_next = ( torch.where( - chunk == 1, - noise, - torch.where( - chunk % 2 < 1, - chunk // 2 if integer_division else chunk / 2, - chunk * 3 + chunk.sign(), - ), + prev_evens, + prev_muls if emul == 1 else prev_muls * emul, + prev_muls if omul == 1 else prev_muls * omul, ) + if emul != 1 or omul != 1 + else prev_muls ) - result = result.sub_(rmin).div_(rmaxsubmin).to(dtype=dtype) - return result[ - tuple(slice(None, sz) for sz in (shape if flatten else out_shape)) - ].reshape(out_shape) + muls[slice_curr] = ( + torch.where(need_reset, 1.0, muls_next) + if need_reset is not False + else muls_next + ) + curr_muls = muls[slice_curr] + prev_adds_scaled = prev_adds * curr_muls + prev_sign = prev.sign() if keepsign else 1.0 + adds_next = ( + torch.where( + prev_evens, + prev_adds_scaled + if eadd == 0 + else prev_adds_scaled + eadd * prev_sign, + prev_adds_scaled + if oadd == 0 + else prev_adds_scaled + oadd * prev_sign, + ) + if eadd != 0 or oadd != 0 + else prev_adds_scaled + ) + adds[slice_curr] = ( + torch.where(need_reset, 0.0, adds_next) + if need_reset is not False + else adds_next + ) + curr_adds = adds[slice_curr] + result_next = utils.maybe_apply( + (noise * curr_muls).add_(curr_adds), + intmode, + torch.trunc, + ) + result[slice_curr] = ( + torch.where(need_reset, noise, result_next) + if need_reset is not False + else result_next + ) + output_slice = tuple( + slice(None, sz) for sz in (shape if flatten else out_shape) + ) + return self._iteration_output( + *args, + result_chains=result, + orig_noise=orig_noise, + noise=noise, + raw_adds=adds, + muls=muls, + chain_length=chain_length, + dim=dim, + output_shape=out_shape, + output_slice=output_slice, + dtype=dtype, + ) - def generate(self, *_args): + def _trim_chain_offset( + self, + t: torch.Tensor, + dim: int, + chain_length: int, + ) -> torch.Tensor: + co = self.chain_offset + if co < 1: + return t + chunks = t.split(chain_length, dim) + slices = [slice(None)] * t.ndim + slices[dim] = slice(co, None) + return torch.cat( + tuple(chunk[slices] for chunk in chunks), + dim=dim, + ) + + def _iteration_output( + self, + *args, + result_chains: torch.Tensor, + orig_noise: torch.Tensor, + noise: torch.Tensor, + raw_adds: torch.Tensor, + muls: torch.Tensor, + chain_length: int, + dim: int, + output_shape: Sequence, + output_slice: Sequence, + dtype: str | torch.dtype, + ) -> torch.Tensor: + omode = self.output_mode + quantile = self.quantile + noise_exp = noise.repeat_interleave(chain_length, dim) + nadds = raw_adds.div_(noise_exp) + ratios = result_chains / noise_exp + if omode in {"values", "ratios", "seed_x_ratios", "noise_x_ratios"}: + out1 = ratios + elif omode in {"mults", "seed_x_mults", "noise_x_mults"}: + out1 = muls + elif omode in {"adds", "seed_x_adds", "noise_x_adds"}: + out1 = nadds + else: + raise ValueError("Bad output mode") + out1 = self._trim_chain_offset(out1, dim=dim, chain_length=chain_length) + if quantile not in {0, 1}: + out1 = utils.quantile_normalize( + out1, + quantile=quantile, + dim=0, + strategy=self.quantile_strategy, + ) + out1 = out1[output_slice].reshape(output_shape).to(dtype=dtype) + if omode in {"ratios", "mults", "adds"}: + return out1 + if omode in {"values", "seed_x_ratios", "seed_x_mults", "seed_x_adds"}: + out2 = orig_noise.repeat_interleave(chain_length - self.chain_offset, dim) + elif omode in {"noise_x_ratios", "noise_x_mults", "noise_x_adds"}: + out2 = ( + self.rand_like(dtype=out1.dtype) + if self.mix_noise_sampler is None + else self.mix_noise_sampler(*args) + ) + out2 = out2[output_slice].reshape(output_shape).to(dtype=dtype) + return out2 * out1 + + def generate(self, *args): out_dims = len(self.shape) dims = tuple(dim if dim >= 0 else out_dims + dim for dim in self.dims) n_dims, n_chainlens = len(dims), len(self.chain_length) @@ -1472,13 +1783,13 @@ class CollatzNoiseGenerator(NoiseGenerator): raise ValueError("Dimension out of range") dtype, device = self.dtype, self.device result = torch.zeros(self.shape, dtype=dtype, device=device) - gen_function = partial( - self._generate_iteration, - integer_division=self.variant == 2, - ) it_scale = 1.0 / self.iterations for iteration in range(self.iterations): - temp = gen_function( + if iteration > 0 and (iteration % 25) == 0: + # It's soooo slow! + throw_exception_if_processing_interrupted() + temp = self._generate_iteration( + *args, dim=dims[iteration % n_dims], chain_length=self.chain_length[iteration % n_chainlens], flatten=self.flatten, @@ -1488,7 +1799,12 @@ class CollatzNoiseGenerator(NoiseGenerator): ) result += temp if self.adjust_scale: - result = normalize_to_scale(result, -1.0, 1.0, dim=1) + result = normalize_to_scale( + result, + -1.0, + 1.0, + dim=tuple(range(1 if result.ndim < 4 else 2, result.ndim)), + ) return result @@ -1512,5 +1828,6 @@ __all__ = ( "PyramidOldNoiseGenerator", "StudentTNoiseGenerator", "UniformNoiseGenerator", + "WaveletFilterNoiseGenerator", "WaveletNoiseGenerator", ) diff --git a/py/utils.py b/py/utils.py index a492fc1..06d2bbb 100644 --- a/py/utils.py +++ b/py/utils.py @@ -16,6 +16,7 @@ UPSCALE_METHODS = ( "area", "bicubic", "bislerp", + "adaptive_avg_pool2d", ) @@ -26,6 +27,8 @@ def scale_samples( *, mode: str = "bicubic", ) -> torch.Tensor: + if mode == "adaptive_avg_pool2d": + return torch.nn.functional.adaptive_avg_pool2d(samples, (height, width)) return common_upscale(samples, width, height, mode, None) @@ -83,6 +86,33 @@ 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) + + +quantile_handlers = { + "clamp": lambda noise, nq: 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( + 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( + 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), +} + + # Initial version based on Studentt distribution normalizatino from https://github.com/Clybius/ComfyUI-Extra-Samplers/ def quantile_normalize( noise: torch.Tensor, @@ -93,6 +123,7 @@ def quantile_normalize( nq_fac: float = 1.0, pow_fac: float = 0.5, strategy: str = "clamp", + strategy_handler=None, ) -> torch.Tensor: if quantile is None or quantile <= 0 or quantile >= 1: return noise @@ -128,20 +159,15 @@ def quantile_normalize( ) nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim) nq = nq.mul_(nq_fac).reshape(*nq_shape) - if strategy == "clamp": - noise = noise.clamp(-nq, nq) - elif strategy == "zero": - noise = torch.where((noise < -nq) | (noise > nq), 0, noise) - elif strategy == "half": - noise = torch.where((noise < -nq) | (noise > nq), noise * 0.5, noise) - elif strategy == "tenth": - noise = torch.where((noise < -nq) | (noise > nq), noise * 0.1, noise) - else: - raise ValueError("Unknown strategy") - noise = torch.copysign( - torch.pow(torch.abs(noise), pow_fac), - noise, + handler = ( + quantile_handlers.get(strategy) + if strategy_handler is None + else strategy_handler ) + if handler is None: + raise ValueError("Unknown strategy") + noise = handler(noise, nq) + 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() @@ -246,3 +272,14 @@ def pattern_break( if restore_scale: noise = normalize_to_scale(noise, orig_min, orig_max, dim=()) return blend_function(noise, result, percentage).to(dtype=orig_dtype) + + +def trunc_decimals(x: torch.Tensor, decimals: int = 3) -> torch.Tensor: + x_i = x.trunc() + x_f = x - x_i + scale = 10.0**decimals + return x_i.add_(x_f.mul_(scale).trunc_().mul_(1.0 / scale)) + + +def maybe_apply(val, cond, fun): + return fun(val) if cond else val