From 78e8451324127ed1013845ed168df2fba17ec7fc Mon Sep 17 00:00:00 2001 From: blepping <157360029+blepping@users.noreply.github.com> Date: Mon, 6 May 2024 16:51:23 -0600 Subject: [PATCH] Feat modulated repeated noise (#5) * Add modulated and repeated noise nodes --- README.md | 25 +++++ __init__.py | 15 +-- changelog.md | 4 + py/nodes.py | 120 ++++++++++++++++++++++- py/noise.py | 269 ++++++++++++++++++++++++++++++++++++++++++++++++++- 5 files changed, 417 insertions(+), 16 deletions(-) diff --git a/README.md b/README.md index 2dc7628..7744b0a 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,27 @@ From a usage perspective, using positive alpha will tend to create a colorful ef Noise from the `SonarCustomNoise` node and `SonarPowerNoise` can be freely mixed. +### `SonarModulatedNoise` + +Experimental noise modulation based on code stolen from +[ComfyUI-Extra-Samplers](https://github.com/Clybius/ComfyUI-Extra-Samplers). _Probably_ does not work correctly +for normal sampling — I expect the modulation will be based on the tensor where the noise sampler was created +rather than each step. However it may be useful for something like restart sampling noise +(see `KRestartSamplerCustomNoise` below). + +*Note*: It's likely this node will be changed in the future. + +### `SonarRepeatedNoise` + +Experimental node to cache noise sampler results. Why would you want to do this? Some noise samplers are +relatively slow (`pyramid` for example) or it may be slow to generate noise if you are mixing many types +of noise. When `permute` is enabled, a random effect like flipping the noise or rolling it in some dimension +will be chosen each time the noise sampler is called. I recommend leaving `permute` on. Note that repeated +noise (especially with `permute` disabled) can be stronger than normal noise, so you may need to rescale to +a value lower than `1.0` or decrease `s_noise` for the sampler. + +*Note*: It's likely this node will be changed in the future. + ### `KRestartSamplerCustomNoise` If you have a recent enough version of [ComfyUI_restart_sampling](https://github.com/ssitu/ComfyUI_restart_sampling/) @@ -105,6 +126,10 @@ installed, you'll also get the `KRestartSamplerCustomNoise` node which is exactl except for adding an optional custom noise input. See the restart sampling repo for more information: https://github.com/ssitu/ComfyUI_restart_sampling +### `RestartSamplerCustomNoise` + +As above, except this is the custom sampler version. + ## Sonar Sampler Parameters Very abbreviated section. The init type can make a big difference. If you use `RANDOM` you can get away with setting `direction` to high values (like up to `2.25` or so) and absurdly low values (like `-30.0`). It's also possible to set `momentum` and `momentum_hist` to negative values, although whether it's a good idea... diff --git a/__init__.py b/__init__.py index 2912708..b473aa0 100644 --- a/__init__.py +++ b/__init__.py @@ -2,20 +2,9 @@ from .py import nodes, powernoise, sonar sonar.add_samplers() -NODE_CLASS_MAPPINGS = { - "SamplerSonarEuler": nodes.SamplerNodeSonarEuler, - "SamplerSonarEulerA": nodes.SamplerNodeSonarEulerAncestral, - "SamplerSonarDPMPPSDE": nodes.SamplerNodeSonarDPMPPSDE, - "SamplerConfigOverride": nodes.SamplerNodeConfigOverride, - "NoisyLatentLike": nodes.NoisyLatentLikeNode, - "SonarCustomNoise": nodes.SonarCustomNoiseNode, +NODE_CLASS_MAPPINGS = nodes.NODE_CLASS_MAPPINGS | { "SonarPowerNoise": powernoise.SonarPowerNoiseNode, - "SonarGuidanceConfig": nodes.GuidanceConfigNode, } - -NODE_DISPLAY_NAME_MAPPINGS = {} - -if hasattr(nodes, "KRestartSamplerCustomNoise"): - NODE_CLASS_MAPPINGS["KRestartSamplerCustomNoise"] = nodes.KRestartSamplerCustomNoise +NODE_DISPLAY_NAME_MAPPINGS = nodes.NODE_DISPLAY_NAME_MAPPINGS __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/changelog.md b/changelog.md index 1f8708a..645d5e3 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,10 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20240506 + +* Add `SonarModulatedNoise` and `SonarRepeatedNoise` nodes. + ## 20240327 * Fixed issue when using Sonar samplers in normal sampling nodes/via stuff like `KSamplerSelect`. diff --git a/py/nodes.py b/py/nodes.py index 577c2ef..b473f18 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -168,6 +168,65 @@ class SonarCustomNoiseNode(SonarCustomNoiseNodeBase): return noise.CustomNoiseItem +class SonarModulatedNoiseNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), + "modulation_type": ( + ( + "intensity", + "frequency", + "spectral_signum", + "none", + ), + ), + "dims": ("INT", {"default": 3, "min": 1, "max": 3}), + "strength": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0}), + }, + } + + RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) + CATEGORY = "advanced/noise" + FUNCTION = "go" + + def go(self, sonar_custom_noise, modulation_type, dims, strength): + return ( + noise.ModulatedNoise( + sonar_custom_noise.make_noise_sampler, + modulation_type=modulation_type, + modulation_strength=strength, + modulation_dims=dims, + ), + ) + + +class SonarRepeatedNoiseNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), + "repeat_length": ("INT", {"default": 8, "min": 1, "max": 100}), + "permute": ("BOOLEAN", {"default": True}), + }, + } + + RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) + CATEGORY = "advanced/noise" + FUNCTION = "go" + + def go(self, sonar_custom_noise, repeat_length, permute=True): + return ( + noise.RepeatedNoise( + sonar_custom_noise.make_noise_sampler, + repeat_length, + permute=permute, + ), + ) + + class GuidanceConfigNode: @classmethod def INPUT_TYPES(cls): @@ -586,6 +645,20 @@ class SamplerNodeConfigOverride: ) +NODE_CLASS_MAPPINGS = { + "SamplerSonarEuler": SamplerNodeSonarEuler, + "SamplerSonarEulerA": SamplerNodeSonarEulerAncestral, + "SamplerSonarDPMPPSDE": SamplerNodeSonarDPMPPSDE, + "SamplerConfigOverride": SamplerNodeConfigOverride, + "NoisyLatentLike": NoisyLatentLikeNode, + "SonarCustomNoise": SonarCustomNoiseNode, + "SonarModulatedNoise": SonarModulatedNoiseNode, + "SonarRepeatedNoise": SonarRepeatedNoiseNode, + "SonarGuidanceConfig": GuidanceConfigNode, +} + +NODE_DISPLAY_NAME_MAPPINGS = {} + try: import custom_nodes.ComfyUI_restart_sampling as rs @@ -597,6 +670,11 @@ try: class KRestartSamplerCustomNoise: @classmethod def INPUT_TYPES(cls): + get_normal_schedulers = getattr( + rs.nodes, + "get_supported_normal_schedulers", + rs.nodes.get_supported_restart_schedulers, + ) return { "required": { "model": ("MODEL",), @@ -608,7 +686,7 @@ try: "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), "sampler": ("SAMPLER",), - "scheduler": (tuple(rs.restart_sampling.SCHEDULER_MAPPING.keys()),), + "scheduler": (get_normal_schedulers(),), "positive": ("CONDITIONING",), "negative": ("CONDITIONING",), "latent_image": ("LATENT",), @@ -676,5 +754,45 @@ try: if custom_noise_opt else None, ) + + NODE_CLASS_MAPPINGS["KRestartSamplerCustomNoise"] = KRestartSamplerCustomNoise + + if not hasattr(rs.restart_sampling, "RestartSampler"): + # Dumb test part II: The Dumbening + raise NotImplementedError # noqa: TRY301 + + class RestartSamplerCustomNoise: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sampler": ("SAMPLER",), + "chunked_mode": ("BOOLEAN", {"default": True}), + }, + "optional": { + "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), + }, + } + + RETURN_TYPES = ("SAMPLER",) + FUNCTION = "go" + CATEGORY = "sampling/custom_sampling/samplers" + + def go(self, sampler, chunked_mode, custom_noise_opt=None): + restart_options = { + "restart_chunked": chunked_mode, + "restart_wrapped_sampler": sampler, + "restart_custom_noise": None + if custom_noise_opt is None + else custom_noise_opt.make_noise_sampler, + } + restart_sampler = samplers.KSAMPLER( + rs.restart_sampling.RestartSampler.sampler_function, + extra_options=sampler.extra_options | restart_options, + inpaint_options=sampler.inpaint_options, + ) + return (restart_sampler,) + + NODE_CLASS_MAPPINGS["RestartSamplerCustomNoise"] = RestartSamplerCustomNoise except (ImportError, NotImplementedError): pass diff --git a/py/noise.py b/py/noise.py index c4a4a79..ec8829e 100644 --- a/py/noise.py +++ b/py/noise.py @@ -11,6 +11,7 @@ from typing import Callable import torch from comfy.k_diffusion import sampling from torch import FloatTensor, Generator, Tensor +from torch.distributions import StudentT # ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311 @@ -405,8 +406,6 @@ def pyramid_noise_like(x, generator=None, device="cpu", discount=0.8): def studentt_noise_like(x): - from torch.distributions import StudentT - noise = StudentT(loc=0, scale=0.2, df=1).rsample(x.size()) s: FloatTensor = torch.quantile(noise.flatten(start_dim=1).abs(), 0.75, dim=-1) s = s.reshape(*s.shape, 1, 1, 1) @@ -552,6 +551,272 @@ class NoiseSampler: return noise +class RepeatedNoise: + def __init__(self, noise_sampler, repeat_length, permute=True): + self.noise_sampler = noise_sampler + self.repeat_length = repeat_length + self.permute = permute + + def clone(self): + return RepeatedNoise(self.noise_sampler, self.repeat_length) + + def make_noise_sampler(self, x, *args, **kwargs): + ns = self.noise_sampler(x, *args, **kwargs) + noise_items = [] + permute_options = 2 + u32_max = 0xFFFF_FFFF + seed = kwargs.get("seed") + if seed is None: + seed = torch.randint( + -u32_max, + u32_max, + (1,), + device="cpu", + dtype=torch.int64, + ).item() + gen = torch.Generator(device="cpu") + gen.manual_seed(seed) + + def noise_sampler(s, sn): + rands = torch.randint( + u32_max, + (4,), + generator=gen, + dtype=torch.uint32, + ).tolist() + if len(noise_items) < self.repeat_length: + idx = len(noise_items) + noise_items.append(ns(s, sn)) + else: + idx = rands[0] % self.repeat_length + noise = noise_items[idx] + if not self.permute: + return noise.clone() + noise_dims = len(noise.shape) + match rands[1] % permute_options: + case 0: + if rands[2] <= u32_max // 10: + # 10% of the time we return the original tensor instead of flipping + noise = noise.clone() + else: + dim = -1 + (rands[2] % (noise_dims + 1)) + noise = torch.flip(noise, (dim,)) + case 1: + dim = rands[2] % noise_dims + count = rands[3] % noise.shape[dim] + noise = torch.roll(noise, count, dims=(dim,)).clone() + return noise + + return noise_sampler + + +# Modulated noise functions copied from https://github.com/Clybius/ComfyUI-Extra-Samplers +# They probably don't work correctly for normal sampling. +class ModulatedNoise: + MODULATION_DIMS = (-3, (-2, -1), (-3, -2, -1)) + + def __init__( + self, + noise_sampler, + modulation_type="none", + modulation_strength=2.0, + modulation_dims=3, + ): + self.noise_sampler = noise_sampler + self.dims = self.MODULATION_DIMS[modulation_dims - 1] + self.type = modulation_type + self.strength = modulation_strength + match self.type: + case "intensity": + self.modulation_function = self.intensity_based_multiplicative_noise + case "frequency": + self.modulation_function = self.frequency_based_noise + case "spectral_signum": + self.modulation_function = self.spectral_modulate_noise + case _: + self.modulation_function = None + + def clone(self): + return ModulatedNoise(self.noise_sampler, self.type, self.strength, self.dims) + + def make_noise_sampler(self, x, *args, **kwargs): + ns = self.noise_sampler(x, *args, **kwargs) + if not self.modulation_function: + return ns + s_noise = sigma_up = 1.0 + return lambda s, sn: self.modulation_function( + x, + ns(s, sn), + s_noise, + sigma_up, + self.strength, + self.dims, + ) + + @staticmethod + def intensity_based_multiplicative_noise( + x, + noise, + s_noise, + sigma_up, + intensity, + dims, + ) -> torch.Tensor: + """Scales noise based on the intensities of the input tensor.""" + std = torch.std( + x - x.mean(), + dim=dims, + keepdim=True, + ) # Average across channels to get intensity + scaling = ( + 1 / (std * abs(intensity) + 1.0) + ) # Scale std by intensity, as not doing this leads to more noise being left over, leading to crusty/preceivably extremely oversharpened images + additive_noise = noise * s_noise * sigma_up + scaled_noise = noise * s_noise * sigma_up * scaling + additive_noise + + noise_norm = torch.norm(additive_noise) + scaled_noise_norm = torch.norm(scaled_noise) + scaled_noise *= noise_norm / scaled_noise_norm # Scale to normal noise strength + return scaled_noise * intensity + additive_noise * (1 - intensity) + + @staticmethod + def frequency_based_noise( + z_k, + noise, + s_noise, + sigma_up, + intensity, + channels, + ) -> torch.Tensor: + """Scales the high-frequency components of the noise based on the given intensity.""" + additive_noise = noise * s_noise * sigma_up + + std = torch.std( + z_k - z_k.mean(), + dim=channels, + keepdim=True, + ) # Average across channels to get intensity + scaling = 1 / (std * abs(intensity) + 1.0) + # Perform Fast Fourier Transform (FFT) + z_k_freq = torch.fft.fft2(scaling * additive_noise + additive_noise) + + # Get the magnitudes of the frequency components + magnitudes = torch.abs(z_k_freq) + + # Create a high-pass filter (emphasize high frequencies) + h, w = z_k.shape[-2:] + b = abs( + intensity, + ) # Controls the emphasis of the high pass (higher frequencies are boosted) + high_pass_filter = 1 - torch.exp( + -((torch.arange(h)[:, None] / h) ** 2 + (torch.arange(w)[None, :] / w) ** 2) + * b**2, + ) + high_pass_filter = high_pass_filter.to(z_k.device) + + # Apply the filter to the magnitudes + magnitudes_scaled = magnitudes * (1 + high_pass_filter) + + # Reconstruct the complex tensor with scaled magnitudes + z_k_freq_scaled = magnitudes_scaled * torch.exp(1j * torch.angle(z_k_freq)) + + # Perform Inverse Fast Fourier Transform (IFFT) + z_k_scaled = torch.fft.ifft2(z_k_freq_scaled) + + # Return the real part of the result + z_k_scaled = torch.real(z_k_scaled) + + noise_norm = torch.norm(additive_noise) + scaled_noise_norm = torch.norm(z_k_scaled) + + z_k_scaled *= noise_norm / scaled_noise_norm # Scale to normal noise strength + + return z_k_scaled * intensity + additive_noise * (1 - intensity) + + @staticmethod + def spectral_modulate_noise( + _unused, + noise, + s_noise, + sigma_up, + intensity, + channels, + spectral_mod_percentile=5.0, + ) -> torch.Tensor: # Modified for soft quantile adjustment using a novel:tm::c::r: method titled linalg. + additive_noise = noise * s_noise * sigma_up + # Convert image to Fourier domain + fourier = torch.fft.fftn( + additive_noise, + dim=channels, + ) # Apply FFT along Height and Width dimensions + + log_amp = torch.log(torch.sqrt(fourier.real**2 + fourier.imag**2)) + + quantile_low = ( + torch.quantile( + log_amp.abs().flatten(1), + spectral_mod_percentile * 0.01, + dim=1, + ) + .unsqueeze(-1) + .unsqueeze(-1) + .expand(log_amp.shape) + ) + + quantile_high = ( + torch.quantile( + log_amp.abs().flatten(1), + 1 - (spectral_mod_percentile * 0.01), + dim=1, + ) + .unsqueeze(-1) + .unsqueeze(-1) + .expand(log_amp.shape) + ) + + quantile_max = ( + torch.quantile(log_amp.abs().flatten(1), 1, dim=1) + .unsqueeze(-1) + .unsqueeze(-1) + .expand(log_amp.shape) + ) + + # Decrease high-frequency components + mask_high = log_amp > quantile_high # If we're larger than 95th percentile + + additive_mult_high = torch.where( + mask_high, + 1 + - ((log_amp - quantile_high) / (quantile_max - quantile_high)).clamp_( + max=0.5, + ), # (1) - (0-1), where 0 is 95th %ile and 1 is 100%ile + torch.tensor(1.0), + ) + + # Increase low-frequency components + mask_low = log_amp < quantile_low + additive_mult_low = torch.where( + mask_low, + 1 + + (1 - (log_amp / quantile_low)).clamp_( + max=0.5, + ), # (1) + (0-1), where 0 is 5th %ile and 1 is 0%ile + torch.tensor(1.0), + ) + + mask_mult = (additive_mult_low * additive_mult_high) ** intensity + # print(mask_mult) + filtered_fourier = fourier * mask_mult + + # Inverse transform back to spatial domain + inverse_transformed = torch.fft.ifftn( + filtered_fourier, + dim=channels, + ) # Apply IFFT along Height and Width dimensions + + return inverse_transformed.real.to(additive_noise.device) + + NOISE_SAMPLERS: dict[NoiseType, Callable] = { NoiseType.BROWNIAN: NoiseSampler.wrap(sampling.BrownianTreeNoiseSampler), NoiseType.GAUSSIAN: NoiseSampler.simple(torch.randn_like),