diff --git a/py/nodes.py b/py/nodes.py index b473f18..19e1c21 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -10,6 +10,7 @@ from comfy import samplers from . import noise from .noise import NoiseType +from .noise_generation import scale_noise from .sonar import ( GuidanceConfig, GuidanceType, @@ -92,8 +93,7 @@ class NoisyLatentLikeNode: result = ns(None, None) finally: torch.random.set_rng_state(randst) - if multiplier != 1.0: - result *= multiplier + result = scale_noise(result, multiplier, normalized=True) if add_to_latent: result += latent_samples.to(result.device) return ({"samples": result},) @@ -173,6 +173,16 @@ class SonarModulatedNoiseNode: def INPUT_TYPES(cls): return { "required": { + "factor": ( + "FLOAT", + { + "default": 1.0, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), "modulation_type": ( ( @@ -184,22 +194,55 @@ class SonarModulatedNoiseNode: ), "dims": ("INT", {"default": 3, "min": 1, "max": 3}), "strength": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0}), + "normalize_result": (("default", "forced", "disabled"),), + "normalize_noise": (("default", "forced", "disabled"),), + "normalize_ref": ( + "BOOLEAN", + {"default": True}, + ), }, + "optional": {"ref_latent_opt": ("LATENT",)}, } RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) CATEGORY = "advanced/noise" FUNCTION = "go" - def go(self, sonar_custom_noise, modulation_type, dims, strength): - return ( + def go( + self, + factor, + sonar_custom_noise, + modulation_type, + dims, + strength, + normalize_result, + normalize_noise, + normalize_ref, + ref_latent_opt=None, + ): + normalize_result = ( + None if normalize_result == "default" else normalize_result == "forced" + ) + normalize_noise = ( + None if normalize_noise == "default" else normalize_noise == "forced" + ) + if ref_latent_opt is not None: + ref_latent_opt = ref_latent_opt["samples"].clone() + nis = noise.CustomNoiseChain() + nis.add( noise.ModulatedNoise( - sonar_custom_noise.make_noise_sampler, + factor, + sonar_custom_noise.rescaled(1.0).make_noise_sampler, + normalize_result, + normalize_noise, + normalize_ref, modulation_type=modulation_type, modulation_strength=strength, modulation_dims=dims, + ref_latent_opt=ref_latent_opt, ), ) + return (nis,) class SonarRepeatedNoiseNode: @@ -207,8 +250,19 @@ class SonarRepeatedNoiseNode: def INPUT_TYPES(cls): return { "required": { + "factor": ( + "FLOAT", + { + "default": 1.0, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), "repeat_length": ("INT", {"default": 8, "min": 1, "max": 100}), + "normalize": (("default", "forced", "disabled"),), "permute": ("BOOLEAN", {"default": True}), }, } @@ -217,14 +271,206 @@ class SonarRepeatedNoiseNode: CATEGORY = "advanced/noise" FUNCTION = "go" - def go(self, sonar_custom_noise, repeat_length, permute=True): - return ( + def go(self, factor, sonar_custom_noise, repeat_length, normalize, permute=True): + normalize = None if normalize == "default" else normalize == "forced" + nis = noise.CustomNoiseChain() + nis.add( noise.RepeatedNoise( - sonar_custom_noise.make_noise_sampler, + factor, + sonar_custom_noise.rescaled(1.0).make_noise_sampler, repeat_length, + normalize, permute=permute, ), ) + return (nis,) + + +class SonarScheduledNoiseNode: + RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) + CATEGORY = "advanced/noise" + FUNCTION = "go" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "factor": ( + "FLOAT", + { + "default": 1.0, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), + "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), + "normalize": (("default", "forced", "disabled"),), + }, + "optional": {"fallback_sonar_custom_noise": ("SONAR_CUSTOM_NOISE",)}, + } + + def go( + self, + model, + factor, + sonar_custom_noise, + start_percent, + end_percent, + normalize, + fallback_sonar_custom_noise=None, + ): + normalize = None if normalize == "default" else normalize == "forced" + ms = model.get_model_object("model_sampling") + start_sigma = ms.percent_to_sigma(start_percent) + end_sigma = ms.percent_to_sigma(end_percent) + return ( + noise.CustomNoiseChain( + [ + noise.ScheduledNoise( + factor, + sonar_custom_noise.rescaled(1.0).make_noise_sampler, + start_sigma, + end_sigma, + normalize, + fallback_noise_sampler=fallback_sonar_custom_noise.rescaled( + 1.0, + ).make_noise_sampler + if fallback_sonar_custom_noise is not None + else None, + ), + ], + ), + ) + + +class SonarCompositeNoiseNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "factor": ( + "FLOAT", + { + "default": 1.0, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), + "sonar_custom_noise_dst": ("SONAR_CUSTOM_NOISE",), + "sonar_custom_noise_src": ("SONAR_CUSTOM_NOISE",), + "normalize_dst": (("default", "forced", "disabled"),), + "normalize_src": (("default", "forced", "disabled"),), + "mask": ("MASK",), + }, + } + + RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) + CATEGORY = "advanced/noise" + FUNCTION = "go" + + def go( + self, + factor, + sonar_custom_noise_dst, + sonar_custom_noise_src, + normalize_src, + normalize_dst, + mask, + ): + normalize_src = ( + None if normalize_src == "default" else normalize_src == "forced" + ) + normalize_dst = ( + None if normalize_dst == "default" else normalize_dst == "forced" + ) + nis = noise.CustomNoiseChain() + + nis.add( + noise.CompositeNoise( + factor, + sonar_custom_noise_dst.rescaled(1.0).make_noise_sampler, + sonar_custom_noise_src.rescaled(1.0).make_noise_sampler, + normalize_src, + normalize_dst, + mask.clone(), + ), + ) + return (nis,) + + +class SonarGuidedNoiseNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + "sonar_custom_noise": ("SONAR_CUSTOM_NOISE",), + "method": (("euler", "linear"),), + "factor": ( + "FLOAT", + { + "default": 1.0, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), + "guidance_factor": ( + "FLOAT", + { + "default": 0.0125, + "min": -100.0, + "max": 100.0, + "step": 0.001, + "round": False, + }, + ), + "normalize": (("default", "forced", "disabled"),), + "normalize_ref": ( + "BOOLEAN", + {"default": True}, + ), + }, + } + + RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) + CATEGORY = "advanced/noise" + FUNCTION = "go" + + def go( + self, + latent, + sonar_custom_noise, + normalize, + normalize_ref=True, + method="euler", + factor=1.0, + guidance_factor=0.5, + ): + from .sonar import SonarGuidanceMixin + + normalize = None if normalize == "default" else normalize == "forced" + nis = noise.CustomNoiseChain() + nis.add( + noise.GuidedNoise( + factor, + guidance_factor, + SonarGuidanceMixin.prepare_ref_latent(latent["samples"].clone()), + sonar_custom_noise.rescaled(1.0).make_noise_sampler, + method, + normalize, + normalize_ref, + ), + ) + return (nis,) class GuidanceConfigNode: @@ -625,6 +871,7 @@ class SamplerNodeConfigOverride: sigma_max, seed=seed, cpu=True, + normalized=True, ) sig = inspect.signature(sampler.sampler_function) params = sig.parameters @@ -652,8 +899,11 @@ NODE_CLASS_MAPPINGS = { "SamplerConfigOverride": SamplerNodeConfigOverride, "NoisyLatentLike": NoisyLatentLikeNode, "SonarCustomNoise": SonarCustomNoiseNode, + "SonarCompositeNoise": SonarCompositeNoiseNode, "SonarModulatedNoise": SonarModulatedNoiseNode, "SonarRepeatedNoise": SonarRepeatedNoiseNode, + "SonarScheduledNoise": SonarScheduledNoiseNode, + "SonarGuidedNoise": SonarGuidedNoiseNode, "SonarGuidanceConfig": GuidanceConfigNode, } diff --git a/py/noise.py b/py/noise.py index ec8829e..24d582a 100644 --- a/py/noise.py +++ b/py/noise.py @@ -1,65 +1,20 @@ -# Noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers from __future__ import annotations import abc import functools as fun -import math import operator as op -from enum import Enum, auto from typing import Callable +import comfy import torch from comfy.k_diffusion import sampling -from torch import FloatTensor, Generator, Tensor -from torch.distributions import StudentT +from torch import Tensor + +from .noise_generation import * # ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311 -def scale_noise(noise, factor=1.0, threshold_std_devs=2.5): - mean, std = noise.mean().item(), noise.std().item() - threshold = threshold_std_devs / math.sqrt(noise.numel()) - if abs(mean) > threshold: - noise -= mean - if abs(1.0 - std) > threshold: - noise /= std - if factor != 1.0: - noise *= factor - return noise - - -class NoiseType(Enum): - GAUSSIAN = auto() - UNIFORM = auto() - BROWNIAN = auto() - PERLIN = auto() - STUDENTT = auto() - HIGHRES_PYRAMID = auto() - PYRAMID = auto() - PINK = auto() - LAPLACIAN = auto() - POWER = auto() - RAINBOW_MILD = auto() - # RAINBOW_MILD2 = auto() - RAINBOW_INTENSE = auto() - # RAINBOW_INTENSE2 = auto() - # RAINBOW_INTENSE3 = auto() - GREEN_TEST = auto() - - @classmethod - def get_names(cls, default=None, skip=None): - if default is not None: - yield default.name.lower() - for nt in cls: - if nt == default or (skip and nt in skip): - continue - yield nt.name.lower() - - -class NoiseError(Exception): - pass - - class CustomNoiseItemBase(abc.ABC): def __init__(self, factor, **kwargs): self.factor = factor @@ -82,6 +37,7 @@ class CustomNoiseItemBase(abc.ABC): sigma_max=None, seed=None, cpu=True, + normalized=True, ): raise NotImplementedError @@ -100,6 +56,7 @@ class CustomNoiseItem(CustomNoiseItemBase): sigma_max=None, seed=None, cpu=True, + normalized=True, ): return get_noise_sampler( self.noise_type, @@ -109,6 +66,7 @@ class CustomNoiseItem(CustomNoiseItemBase): seed=seed, cpu=cpu, factor=self.factor, + normalized=normalized, ) @@ -122,15 +80,22 @@ class CustomNoiseChain: ) def add(self, item): + if item is None: + raise ValueError("Attempt to add nil item") self.items.append(item) + @property + def factor(self): + return sum(abs(i.factor) for i in self.items) + def rescaled(self, scale=1.0): - total = sum(i.factor for i in self.items) - divisor = total / scale + divisor = self.factor / scale divisor = divisor if divisor != 0 else 1.0 - return CustomNoiseChain( - [i.clone().set_factor(i.factor / divisor) for i in self.items], - ) + result = self.clone() + if divisor != 1: + for i in result.items: + i.set_factor(i.factor / divisor) + return result @torch.no_grad() def make_noise_sampler( @@ -140,6 +105,7 @@ class CustomNoiseChain: sigma_max=None, seed=None, cpu=True, + normalized=True, ) -> Callable: noise_samplers = tuple( i.make_noise_sampler( @@ -148,349 +114,26 @@ class CustomNoiseChain: sigma_max, seed=seed, cpu=cpu, + normalized=False, ) for i in self.items ) if not noise_samplers or not all(noise_samplers): raise ValueError("Failed to get noise sampler") - scale = sum(i.factor for i in self.items) + factor = self.factor def noise_sampler(sigma, sigma_next): result = fun.reduce( op.add, (ns(sigma, sigma_next) for ns in noise_samplers), ) - return scale_noise(result, scale) + if normalized: + return scale_noise(result, factor) + return result.mul_(factor) return noise_sampler -def get_positions(block_shape: tuple[int, int]) -> Tensor: - """ - Generate position tensor. - - Arguments: - block_shape -- (height, width) of position tensor - - Returns: - position vector shaped (1, height, width, 1, 1, 2) - """ - bh, bw = block_shape - return torch.stack( - torch.meshgrid( - [(torch.arange(b) + 0.5) / b for b in (bw, bh)], - indexing="xy", - ), - -1, - ).view(1, bh, bw, 1, 1, 2) - - -def unfold_grid(vectors: Tensor) -> Tensor: - """ - Unfold vector grid to batched vectors. - - Arguments: - vectors -- grid vectors - - Returns: - batched grid vectors - """ - batch_size, _, gpy, gpx = vectors.shape - return ( - torch.nn.functional.unfold(vectors, (2, 2)) - .view(batch_size, 2, 4, -1) - .permute(0, 2, 3, 1) - .view(batch_size, 4, gpy - 1, gpx - 1, 2) - ) - - -def smooth_step(t: Tensor) -> Tensor: - """ - Smooth step function [0, 1] -> [0, 1]. - - Arguments: - t -- input values (any shape) - - Returns: - output values (same shape as input values) - """ - return t * t * (3.0 - 2.0 * t) - - -def perlin_noise_tensor( - vectors: Tensor, - positions: Tensor, - step: Callable | None = None, -) -> Tensor: - """ - Generate perlin noise from batched vectors and positions. - - Arguments: - vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2) - positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2) - - Keyword Arguments: - step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`) - - Raises: - Exception: if position and vector shapes do not match - - Returns: - (batch_size, block_height * grid_height, block_width * grid_width) - """ - if step is None: - step = smooth_step - - batch_size = vectors.shape[0] - # grid height, grid width - gh, gw = vectors.shape[2:4] - # block height, block width - bh, bw = positions.shape[1:3] - - for i in range(2): - if positions.shape[i + 3] not in (1, vectors.shape[i + 2]): - msg = f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})" - raise NoiseError(msg) - - if positions.shape[0] not in (1, batch_size): - msg = f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})" - raise NoiseError(msg) - - vectors = vectors.view(batch_size, 4, 1, gh * gw, 2) - positions = positions.view(positions.shape[0], bh * bw, -1, 2) - - step_x = step(positions[..., 0]) - step_y = step(positions[..., 1]) - - row0 = torch.lerp( - (vectors[:, 0] * positions).sum(dim=-1), - (vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1), - step_x, - ) - row1 = torch.lerp( - (vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1), - (vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1), - step_x, - ) - noise = torch.lerp(row0, row1, step_y) - return ( - noise.view( - batch_size, - bh, - bw, - gh, - gw, - ) - .permute(0, 3, 1, 4, 2) - .reshape(batch_size, gh * bh, gw * bw) - ) - - -def perlin_noise( - grid_shape: tuple[int, int], - out_shape: tuple[int, int], - batch_size: int = 1, - generator: Generator | None = None, - *args, - **kwargs, -) -> Tensor: - """ - Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation. - - Arguments: - grid_shape -- Shape of grid (height, width). - out_shape -- Shape of output noise image (height, width). - - Keyword Arguments: - batch_size -- (default: {1}) - generator -- random generator used for grid vectors (default: {None}) - - Raises: - Exception: if grid and out shapes do not match - - Returns: - Noise image shaped (batch_size, height, width) - """ - # grid height and width - gh, gw = grid_shape - # output height and width - oh, ow = out_shape - # block height and width - bh, bw = oh // gh, ow // gw - - if oh != bh * gh: - msg = f"Output height {oh} must be divisible by grid height {gh}" - raise NoiseError(msg) - if ow != bw * gw != 0: - msg = f"Output width {ow} must be divisible by grid width {gw}" - raise NoiseError(msg) - - angle = torch.empty( - [batch_size] + [s + 1 for s in grid_shape], - *args, - **kwargs, - ).uniform_(to=2.0 * math.pi, generator=generator) - # random vectors on grid points - vectors = unfold_grid(torch.stack((torch.cos(angle), torch.sin(angle)), dim=1)) - # positions inside grid cells [0, 1) - positions = get_positions((bh, bw)).to(vectors) - return perlin_noise_tensor(vectors, positions).squeeze(0) - - -def rand_perlin_like(x): - noise = torch.randn_like(x) / 2.0 - noise_height = noise.size(dim=2) - noise_width = noise.size(dim=3) - for _ in range(2): - noise += perlin_noise( - (noise_height, noise_width), - (noise_height, noise_width), - batch_size=x.shape[1], # This should be the number of channels. - ).to(x.device) - return noise / noise.std() - - -def uniform_noise_like(x): - return (torch.rand_like(x) - 0.5) * 3.46 - - -def highres_pyramid_noise_like(x, discount=0.7): - ( - b, - c, - h, - w, - ) = x.shape # EDIT: w and h get over-written, rename for a different variant! - orig_h = h - orig_w = w - u = torch.nn.Upsample(size=(orig_h, orig_w), mode="bilinear") - noise = uniform_noise_like(x) - rs = torch.rand(4, dtype=torch.float32) * 2 + 2 - for i in range(4): - r = rs[i] - h, w = min(orig_h * 15, int(h * (r**i))), min(orig_w * 15, int(w * (r**i))) - noise += u(torch.randn(b, c, h, w).to(x)) * discount**i - if h >= orig_h * 15 or w >= orig_w * 15: - break # Lowest resolution is 1x1 - return noise / noise.std() # Scaled back to roughly unit variance - - -def pyramid_noise_like(x, generator=None, device="cpu", discount=0.8): - size = x.size() - b, c, h, w = size - orig_h = h - orig_w = w - noise = torch.zeros(size=size, dtype=x.dtype, layout=x.layout, device=device) - r = 1 - for i in range(5): - r *= 2 # Rather than always going 2x, - noise += ( - torch.nn.functional.interpolate( - ( - torch.normal( - mean=0, - std=0.5**i, - size=(b, c, h * r, w * r), - dtype=x.dtype, - layout=x.layout, - generator=generator, - device=device, - ) - ), - size=(orig_h, orig_w), - mode="nearest-exact", - ) - * discount**i - ) - return noise.to(device=x.device) - - -def studentt_noise_like(x): - 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) - noise = noise.clamp(-s, s) - return torch.copysign(torch.pow(torch.abs(noise), 0.5), noise) - - -def studentt_noise_sampler( - x, -): # Produces more subject-focused outputs due to distribution, unsure if this works - noise = studentt_noise_like(x) - return lambda _sigma, _sigma_next: noise.to(x.device) / (7 / 3) - - -def green_noise_like(x): - # The comments said this didn't work and I had to learn the hard way. Turns out it's true! - width, height = x.size(dim=2), x.size(dim=3) - noise = torch.randn_like(x) - scale = 1.0 / (width * height) - fy = torch.fft.fftfreq(width, device=x.device)[:, None] ** 2 - fx = torch.fft.fftfreq(height, device=x.device) ** 2 - f = fy + fx - power = torch.sqrt(f) - power[0, 0] = 1 - noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power)) - noise *= scale / noise.std() - noise = torch.real(noise).to(x.device) - return noise / noise.std() - - -def generate_1f_noise(tensor, alpha, k, generator=None): - """Generate 1/f noise for a given tensor. - - Args: - tensor: The tensor to add noise to. - alpha: The parameter that determines the slope of the spectrum. - k: A constant. - - Returns: - A tensor with the same shape as `tensor` containing 1/f noise. - """ - fft = torch.fft.fft2(tensor) - freq = torch.arange(1, len(fft) + 1, dtype=torch.float) - spectral_density = k / freq**alpha - return torch.randn(tensor.shape, generator=generator) * spectral_density - - -def pink_noise_like(x): - noise = generate_1f_noise(x, 2.0, 1.0) - noise_mean = torch.mean(noise) - noise_std = torch.std(noise) - return noise.sub_(noise_mean).div_(noise_std).to(x.device) - - -def laplacian_noise_like(x): - from torch.distributions import Laplace - - noise = torch.randn_like(x) / 4.0 - noise += Laplace(loc=0, scale=1.0).rsample(x.size()).to(noise.device) - return noise / noise.std() - - -def power_noise_like(tensor, alpha=2, k=1): # This doesn't work properly right now - """Generate 1/f noise for a given tensor. - - Args: - tensor: The tensor to add noise to. - alpha: The parameter that determines the slope of the spectrum. - k: A constant. - - Returns: - A tensor with the same shape as `tensor` containing 1/f noise. - """ - tensor = torch.randn_like(tensor) - fft = torch.fft.fft2(tensor) - freq = torch.arange(1, len(fft) + 1, dtype=torch.float).reshape( - (len(fft),) + (1,) * (tensor.dim() - 1), - ) - spectral_density = k / freq**alpha - noise = torch.rand(tensor.shape) * spectral_density - mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device) - std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device) - return noise.to(tensor.device).sub_(mean).div_(std) - - class NoiseSampler: def __init__( self, @@ -501,7 +144,7 @@ class NoiseSampler: cpu: bool = False, transform: Callable = lambda t: t, make_noise_sampler: Callable | None = None, - normalize_noise=False, + normalized=False, factor: float = 1.0, ): try: @@ -519,7 +162,7 @@ class NoiseSampler: except TypeError: self.noise_sampler = make_noise_sampler(x) self.factor = factor - self.normalize_noise = normalize_noise + self.normalized = normalized self.transform = transform self.device = x.device self.dtype = x.dtype @@ -543,7 +186,7 @@ class NoiseSampler: noise = self.noise_sampler(*args, **kwargs) noise = ( scale_noise(noise, self.factor) - if self.normalize_noise + if self.normalized else noise.mul_(self.factor) ) if hasattr(noise, "to"): @@ -551,17 +194,215 @@ class NoiseSampler: return noise +class CompositeNoise: + def __init__(self, factor, dst, src, normalize_src, normalize_dst, mask): + self.factor = factor + self.dst_noise_sampler = dst + self.src_noise_sampler = src + self.normalize_src = normalize_src + self.normalize_dst = normalize_dst + self.mask = mask + + def clone(self): + return CompositeNoise( + self.factor, + self.dst_noise_sampler, + self.src_noise_sampler, + self.normalize_src, + self.normalize_dst, + self.mask.clone(), + ) + + def set_factor(self, factor): + self.factor = factor + return self + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + normalize_src = ( + self.normalize_src if self.normalize_src is not None else normalized + ) + normalize_dst = ( + self.normalize_dst if self.normalize_dst is not None else normalized + ) + nsd = self.dst_noise_sampler(x, *args, normalized=False, **kwargs) + nss = self.src_noise_sampler(x, *args, normalized=False, **kwargs) + mask = self.mask.to(x.device, copy=True) + mask = torch.nn.functional.interpolate( + mask.reshape((-1, 1, *mask.shape[-2:])), + size=x.shape[-2:], + mode="bilinear", + ) + mask = comfy.utils.repeat_to_batch_size(mask, x.shape[0]) + imask = torch.ones_like(mask) - mask + + def noise_sampler(s, sn): + noise_dst = scale_noise( + nsd(s, sn), + self.factor, + normalized=normalize_dst, + ).mul_( + imask, + ) + noise_src = scale_noise( + nss(s, sn), + self.factor, + normalized=normalize_src, + ).mul_(mask) + return noise_dst.add_(noise_src) + + return noise_sampler + + +class GuidedNoise: + def __init__( + self, + factor, + guidance_factor, + ref_latent, + noise_sampler, + method, + normalize, + normalize_ref, + ): + self.factor = factor + self.normalize = normalize + self.normalize_ref = normalize_ref + self.ref_latent = ref_latent + self.noise_sampler = noise_sampler + self.method = method + self.guidance_factor = guidance_factor + + def clone(self): + return GuidedNoise( + self.factor, + self.guidance_factor, + self.ref_latent.clone(), + self.noise_sampler, + self.method, + self.normalize, + self.normalize_ref, + ) + + def set_factor(self, factor): + self.factor = factor + return self + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + from .sonar import SonarGuidanceMixin + + normalize = self.normalize if self.normalize is not None else normalized + ns = self.noise_sampler(x, *args, normalized=False, **kwargs) + ref_latent = scale_noise( + self.ref_latent.to(x, copy=True), + normalized=self.normalize_ref, + ) + match self.method: + case "linear": + + def noise_sampler(s, sn): + return scale_noise( + SonarGuidanceMixin.guidance_linear( + scale_noise(ns(s, sn), normalized=normalize), + ref_latent, + self.guidance_factor, + ), + self.factor, + normalized=normalize, + ) + case "euler": + + def noise_sampler(s, sn): + return scale_noise( + SonarGuidanceMixin.guidance_euler( + s, + sn, + scale_noise(ns(s, sn), normalized=normalize), + x, + ref_latent, + self.guidance_factor, + ), + self.factor, + normalized=normalize, + ) + + return noise_sampler + + +class ScheduledNoise: + def __init__( + self, + factor, + noise_sampler, + start_sigma, + end_sigma, + normalize, + fallback_noise_sampler=None, + ): + self.factor = factor + self.noise_sampler = noise_sampler + self.start_sigma = start_sigma + self.end_sigma = end_sigma + self.normalize = normalize + self.fallback_noise_sampler = fallback_noise_sampler + + def clone(self): + return ScheduledNoise( + self.factor, + self.noise_sampler, + self.start_sigma, + self.end_sigma, + self.normalize, + fallback_noise_sampler=self.fallback_noise_sampler, + ) + + def set_factor(self, factor): + self.factor = factor + return self + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + normalize = self.normalize if self.normalize is not None else normalized + ns = self.noise_sampler(x, *args, normalized=False, **kwargs) + if self.fallback_noise_sampler: + nsa = self.fallback_noise_sampler(x, *args, normalized=False, **kwargs) + else: + + def nsa(_s, _sn): + return torch.zeros_like(x) + + def noise_sampler(s, sn): + if s <= self.start_sigma and s >= self.end_sigma: + noise = ns(s, sn) + else: + noise = nsa(s, sn) + return scale_noise(noise, self.factor, normalized=normalize) + + return noise_sampler + + class RepeatedNoise: - def __init__(self, noise_sampler, repeat_length, permute=True): + def __init__(self, factor, noise_sampler, repeat_length, normalize, permute=True): + self.factor = factor + self.normalize = normalize self.noise_sampler = noise_sampler self.repeat_length = repeat_length self.permute = permute def clone(self): - return RepeatedNoise(self.noise_sampler, self.repeat_length) + return RepeatedNoise( + self.factor, + self.noise_sampler, + self.repeat_length, + self.normalize, + self.permute, + ) - def make_noise_sampler(self, x, *args, **kwargs): - ns = self.noise_sampler(x, *args, **kwargs) + def set_factor(self, factor): + self.factor = factor + return self + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + normalize = self.normalize if self.normalize is not None else normalized + ns = self.noise_sampler(x, *args, normalized=False, **kwargs) noise_items = [] permute_options = 2 u32_max = 0xFFFF_FFFF @@ -605,7 +446,7 @@ class RepeatedNoise: dim = rands[2] % noise_dims count = rands[3] % noise.shape[dim] noise = torch.roll(noise, count, dims=(dim,)).clone() - return noise + return scale_noise(noise, self.factor, normalized=normalize) return noise_sampler @@ -617,15 +458,25 @@ class ModulatedNoise: def __init__( self, + factor, noise_sampler, + normalize_result, + normalize_noise, + normalize_ref, modulation_type="none", modulation_strength=2.0, modulation_dims=3, + ref_latent_opt=None, ): + self.factor = factor + self.normalize_result = normalize_result + self.normalize_noise = normalize_noise + self.normalize_ref = normalize_ref self.noise_sampler = noise_sampler - self.dims = self.MODULATION_DIMS[modulation_dims - 1] + self.modulation_dims = modulation_dims self.type = modulation_type self.strength = modulation_strength + self.ref_latent_opt = ref_latent_opt match self.type: case "intensity": self.modulation_function = self.intensity_based_multiplicative_noise @@ -637,21 +488,61 @@ class ModulatedNoise: self.modulation_function = None def clone(self): - return ModulatedNoise(self.noise_sampler, self.type, self.strength, self.dims) + return ModulatedNoise( + self.factor, + self.noise_sampler, + self.normalize_result, + self.normalize_noise, + self.normalize_ref, + self.type, + self.strength, + self.modulation_dims, + self.ref_latent_opt, + ) - def make_noise_sampler(self, x, *args, **kwargs): + def set_factor(self, factor): + self.factor = factor + return self + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + normalize_result = ( + self.normalize_result if self.normalize_result is not None else normalized + ) + normalize_noise = ( + self.normalize_noise if self.normalize_noise is not None else normalized + ) + dims = self.MODULATION_DIMS[self.modulation_dims - 1] 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, - ) + + def noise_sampler(s, sn): + return scale_noise( + ns(s, sn), + self.factor, + normalized=normalize_result or normalize_noise, + ) + + return noise_sampler + + ref_latent = None + if self.ref_latent_opt is not None: + ref_latent = self.ref_latent_opt.to(x, copy=True) + + def noise_sampler(s, sn): + noise = self.modulation_function( + scale_noise( + x if self.ref_latent_opt is None else ref_latent, + normalized=self.normalize_ref, + ), + scale_noise(ns(s, sn), normalized=normalize_noise), + 1.0, # s_noise + 1.0, # sigma_up + self.strength, + dims, + ) + return scale_noise(noise, self.factor, normalized=normalize_result) + + return noise_sampler @staticmethod def intensity_based_multiplicative_noise( @@ -805,7 +696,6 @@ class ModulatedNoise: ) mask_mult = (additive_mult_low * additive_mult_high) ** intensity - # print(mask_mult) filtered_fourier = fourier * mask_mult # Inverse transform back to spatial domain @@ -835,6 +725,25 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = { NoiseType.LAPLACIAN: NoiseSampler.simple(laplacian_noise_like), NoiseType.POWER: NoiseSampler.simple(power_noise_like), NoiseType.GREEN_TEST: NoiseSampler.simple(green_noise_like), + NoiseType.PYRAMID_OLD: NoiseSampler.simple(pyramid_old_noise_like), + NoiseType.PYRAMID_BISLERP: NoiseSampler.simple( + lambda x: pyramid_noise_like(x, upscale_mode="bislerp"), + ), + NoiseType.HIGHRES_PYRAMID_BISLERP: NoiseSampler.simple( + lambda x: highres_pyramid_noise_like(x, upscale_mode="bislerp"), + ), + NoiseType.PYRAMID_AREA: NoiseSampler.simple( + lambda x: pyramid_noise_like(x, upscale_mode="area"), + ), + NoiseType.HIGHRES_PYRAMID_AREA: NoiseSampler.simple( + lambda x: highres_pyramid_noise_like(x, upscale_mode="area"), + ), + NoiseType.PYRAMID_OLD_BISLERP: NoiseSampler.simple( + lambda x: pyramid_old_noise_like(x, upscale_mode="bislerp"), + ), + NoiseType.PYRAMID_OLD_AREA: NoiseSampler.simple( + lambda x: pyramid_old_noise_like(x, upscale_mode="area"), + ), } @@ -846,7 +755,7 @@ def get_noise_sampler( seed: int | None = None, cpu: bool = True, factor: float = 1.0, - normalize_noise=True, + normalized=False, ) -> Callable: if noise_type is None: noise_type = NoiseType.GAUSSIAN @@ -864,5 +773,5 @@ def get_noise_sampler( seed=seed, cpu=cpu, factor=factor, - normalize_noise=normalize_noise, + normalized=normalized, ) diff --git a/py/noise_generation.py b/py/noise_generation.py new file mode 100644 index 0000000..a262c53 --- /dev/null +++ b/py/noise_generation.py @@ -0,0 +1,441 @@ +# Noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers +from __future__ import annotations + +import math +from enum import Enum, auto +from typing import Callable + +import torch +from comfy.utils import common_upscale +from torch import FloatTensor, Generator, Tensor +from torch.distributions import Laplace, StudentT + +# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311 + + +class NoiseType(Enum): + GAUSSIAN = auto() + UNIFORM = auto() + BROWNIAN = auto() + PERLIN = auto() + STUDENTT = auto() + HIGHRES_PYRAMID = auto() + PYRAMID = auto() + PINK = auto() + LAPLACIAN = auto() + POWER = auto() + RAINBOW_MILD = auto() + RAINBOW_INTENSE = auto() + GREEN_TEST = auto() + PYRAMID_OLD = auto() + PYRAMID_BISLERP = auto() + HIGHRES_PYRAMID_BISLERP = auto() + PYRAMID_OLD_BISLERP = auto() + PYRAMID_OLD_AREA = auto() + PYRAMID_AREA = auto() + HIGHRES_PYRAMID_AREA = auto() + + @classmethod + def get_names(cls, default=None, skip=None): + if default is not None: + yield default.name.lower() + for nt in cls: + if nt == default or (skip and nt in skip): + continue + yield nt.name.lower() + + +class NoiseError(Exception): + pass + + +def _scale_noise(noise, factor=1.0, threshold_std_devs=2.5): + if factor != 1.0: + noise *= factor + mean, std = noise.mean().item(), noise.std().item() + threshold = threshold_std_devs / math.sqrt(noise.numel()) + print(f"SCALE: mean={mean}, std={std}") + if abs(mean) > threshold: + noise -= mean + if abs(factor - std) > abs(threshold * factor): + noise /= std + return noise + + +def scale_noise(noise, factor=1.0, normalized=True, threshold_std_devs=2.5): + if not normalized: + return noise.mul_(factor) + mean, std = noise.mean().item(), noise.std().item() + threshold = threshold_std_devs / math.sqrt(noise.numel()) + if abs(mean) > threshold: + noise -= mean + if abs(1.0 - std) > threshold: + noise /= std + if factor != 1.0: + noise *= factor + return noise + + +def get_positions(block_shape: tuple[int, int]) -> Tensor: + """ + Generate position tensor. + + Arguments: + block_shape -- (height, width) of position tensor + + Returns: + position vector shaped (1, height, width, 1, 1, 2) + """ + bh, bw = block_shape + return torch.stack( + torch.meshgrid( + [(torch.arange(b) + 0.5) / b for b in (bw, bh)], + indexing="xy", + ), + -1, + ).view(1, bh, bw, 1, 1, 2) + + +def unfold_grid(vectors: Tensor) -> Tensor: + """ + Unfold vector grid to batched vectors. + + Arguments: + vectors -- grid vectors + + Returns: + batched grid vectors + """ + batch_size, _, gpy, gpx = vectors.shape + return ( + torch.nn.functional.unfold(vectors, (2, 2)) + .view(batch_size, 2, 4, -1) + .permute(0, 2, 3, 1) + .view(batch_size, 4, gpy - 1, gpx - 1, 2) + ) + + +def smooth_step(t: Tensor) -> Tensor: + """ + Smooth step function [0, 1] -> [0, 1]. + + Arguments: + t -- input values (any shape) + + Returns: + output values (same shape as input values) + """ + return t * t * (3.0 - 2.0 * t) + + +def perlin_noise_tensor( + vectors: Tensor, + positions: Tensor, + step: Callable | None = None, +) -> Tensor: + """ + Generate perlin noise from batched vectors and positions. + + Arguments: + vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2) + positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2) + + Keyword Arguments: + step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`) + + Raises: + Exception: if position and vector shapes do not match + + Returns: + (batch_size, block_height * grid_height, block_width * grid_width) + """ + if step is None: + step = smooth_step + + batch_size = vectors.shape[0] + # grid height, grid width + gh, gw = vectors.shape[2:4] + # block height, block width + bh, bw = positions.shape[1:3] + + for i in range(2): + if positions.shape[i + 3] not in (1, vectors.shape[i + 2]): + msg = f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})" + raise NoiseError(msg) + + if positions.shape[0] not in (1, batch_size): + msg = f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})" + raise NoiseError(msg) + + vectors = vectors.view(batch_size, 4, 1, gh * gw, 2) + positions = positions.view(positions.shape[0], bh * bw, -1, 2) + + step_x = step(positions[..., 0]) + step_y = step(positions[..., 1]) + + row0 = torch.lerp( + (vectors[:, 0] * positions).sum(dim=-1), + (vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1), + step_x, + ) + row1 = torch.lerp( + (vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1), + (vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1), + step_x, + ) + noise = torch.lerp(row0, row1, step_y) + return ( + noise.view( + batch_size, + bh, + bw, + gh, + gw, + ) + .permute(0, 3, 1, 4, 2) + .reshape(batch_size, gh * bh, gw * bw) + ) + + +def perlin_noise( + grid_shape: tuple[int, int], + out_shape: tuple[int, int], + batch_size: int = 1, + generator: Generator | None = None, + *args, + **kwargs, +) -> Tensor: + """ + Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation. + + Arguments: + grid_shape -- Shape of grid (height, width). + out_shape -- Shape of output noise image (height, width). + + Keyword Arguments: + batch_size -- (default: {1}) + generator -- random generator used for grid vectors (default: {None}) + + Raises: + Exception: if grid and out shapes do not match + + Returns: + Noise image shaped (batch_size, height, width) + """ + # grid height and width + gh, gw = grid_shape + # output height and width + oh, ow = out_shape + # block height and width + bh, bw = oh // gh, ow // gw + + if oh != bh * gh: + msg = f"Output height {oh} must be divisible by grid height {gh}" + raise NoiseError(msg) + if ow != bw * gw != 0: + msg = f"Output width {ow} must be divisible by grid width {gw}" + raise NoiseError(msg) + + angle = torch.empty( + [batch_size] + [s + 1 for s in grid_shape], + *args, + **kwargs, + ).uniform_(to=2.0 * math.pi, generator=generator) + # random vectors on grid points + vectors = unfold_grid(torch.stack((torch.cos(angle), torch.sin(angle)), dim=1)) + # positions inside grid cells [0, 1) + positions = get_positions((bh, bw)).to(vectors) + return perlin_noise_tensor(vectors, positions).squeeze(0) + + +def rand_perlin_like(x): + noise = torch.randn_like(x) / 2.0 + noise_height = noise.size(dim=2) + noise_width = noise.size(dim=3) + for _ in range(2): + noise += perlin_noise( + (noise_height, noise_width), + (noise_height, noise_width), + batch_size=x.shape[1], # This should be the number of channels. + ).to(x.device) + return scale_noise(noise) + + +def uniform_noise_like(x): + return (torch.rand_like(x) - 0.5) * 3.46 + + +def highres_pyramid_noise_like(x, discount=0.7, upscale_mode="bilinear"): + ( + b, + c, + h, + w, + ) = x.shape # EDIT: w and h get over-written, rename for a different variant! + orig_w, orig_h = w, h + noise = uniform_noise_like(x) + rs = torch.rand(4, dtype=torch.float32) * 2 + 2 + for i in range(4): + r = rs[i] + h, w = min(orig_h * 15, int(h * (r**i))), min(orig_w * 15, int(w * (r**i))) + noise += common_upscale( + torch.randn(b, c, h, w).to(x), + orig_w, + orig_h, + upscale_mode, + None, + ).mul_(discount**i) + if h >= orig_h * 15 or w >= orig_w * 15: + break # Lowest resolution is 1x1 + return scale_noise(noise) + + +def pyramid_old_noise_like( + x, + generator=None, + device="cpu", + discount=0.8, + upscale_mode="nearest-exact", +): + size = x.size() + b, c, h, w = size + orig_h, orig_w = h, w + noise = torch.zeros(size=size, dtype=x.dtype, layout=x.layout, device=device) + r = 1 + for i in range(5): + r *= 2 + noise += common_upscale( + torch.normal( + mean=0, + std=0.5**i, + size=(b, c, h * r, w * r), + dtype=x.dtype, + layout=x.layout, + generator=generator, + device=device, + ), + orig_w, + orig_h, + upscale_mode, + None, + ).mul_(discount**i) + return noise.to(device=x.device) + + +# Copied from https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2 +def pyramid_noise_like(x, discount=0.7, upscale_mode="bilinear"): + b, c, w, h = ( + x.shape + ) # NOTE: w and h get over-written, rename for a different variant! + orig_w, orig_h = w, h + noise = torch.randn_like(x) + for i in range(10): + r = torch.rand(1, device="cpu").item() * 2 + 2 # Rather than always going 2x, + w, h = max(1, int(w / (r**i))), max(1, int(h / (r**i))) + noise += common_upscale( + torch.randn(b, c, w, h).to(x), + orig_h, + orig_w, + upscale_mode, + None, + ).mul_( + discount**i, + ) + if w == 1 or h == 1: + break # Lowest resolution is 1x1 + return scale_noise(noise) + + +def studentt_noise_like(x): + 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) + noise = noise.clamp(-s, s) + return torch.copysign(torch.pow(torch.abs(noise), 0.5), noise) + + +def green_noise_like(x): + # The comments said this didn't work and I had to learn the hard way. Turns out it's true! + width, height = x.size(dim=2), x.size(dim=3) + noise = torch.randn_like(x) + scale = 1.0 / (width * height) + fy = torch.fft.fftfreq(width, device=x.device)[:, None] ** 2 + fx = torch.fft.fftfreq(height, device=x.device) ** 2 + f = fy + fx + power = torch.sqrt(f) + power[0, 0] = 1 + noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power)) + noise *= scale / noise.std() + noise = torch.real(noise).to(x.device) + return scale_noise(noise) + + +def generate_1f_noise(tensor, alpha, k, generator=None): + """Generate 1/f noise for a given tensor. + + Args: + tensor: The tensor to add noise to. + alpha: The parameter that determines the slope of the spectrum. + k: A constant. + + Returns: + A tensor with the same shape as `tensor` containing 1/f noise. + """ + fft = torch.fft.fft2(tensor) + freq = torch.arange(1, len(fft) + 1, dtype=torch.float) + spectral_density = k / freq**alpha + return torch.randn(tensor.shape, generator=generator) * spectral_density + + +def pink_noise_like(x): + return scale_noise(generate_1f_noise(x, 2.0, 1.0)).to(x.device) + # noise = generate_1f_noise(x, 2.0, 1.0) + # noise_mean = torch.mean(noise) + # noise_std = torch.std(noise) + # return noise.sub_(noise_mean).div_(noise_std).to(x.device) + + +def laplacian_noise_like(x): + noise = torch.randn_like(x).div_(4.0) + noise += Laplace(loc=0, scale=1.0).rsample(x.size()).to(noise.device) + return scale_noise(noise) + + +def power_noise_like(tensor, alpha=2, k=1): # This doesn't work properly right now + """Generate 1/f noise for a given tensor. + + Args: + tensor: The tensor to add noise to. + alpha: The parameter that determines the slope of the spectrum. + k: A constant. + + Returns: + A tensor with the same shape as `tensor` containing 1/f noise. + """ + tensor = torch.randn_like(tensor) + fft = torch.fft.fft2(tensor) + freq = torch.arange(1, len(fft) + 1, dtype=torch.float).reshape( + (len(fft),) + (1,) * (tensor.dim() - 1), + ) + spectral_density = k / freq**alpha + noise = torch.rand(tensor.shape).mul_(spectral_density) + mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device) + std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device) + return noise.to(tensor.device).sub_(mean).div_(std) + + +__all__ = ( + "NoiseType", + "NoiseError", + "scale_noise", + "green_noise_like", + "highres_pyramid_noise_like", + "laplacian_noise_like", + "pink_noise_like", + "power_noise_like", + "pyramid_noise_like", + "pyramid_old_noise_like", + "rand_perlin_like", + "studentt_noise_like", + "uniform_noise_like", +) diff --git a/py/powernoise.py b/py/powernoise.py index 298af8c..1dbe29d 100644 --- a/py/powernoise.py +++ b/py/powernoise.py @@ -12,6 +12,7 @@ from torch import Tensor from .nodes import SonarCustomNoiseNodeBase from .noise import CustomNoiseItemBase +from .noise_generation import scale_noise # ruff: noqa: ANN003, FBT001, FBT002 @@ -110,6 +111,7 @@ class PowerNoiseItem(CustomNoiseItemBase): sigma_max: float | None, seed: int | None, cpu: bool = True, + normalized=True, ): shape = x.shape device = x.device @@ -159,7 +161,7 @@ class PowerNoiseItem(CustomNoiseItemBase): if common_mode > 0.0: noise = channel_mixer @ noise.swapaxes(0, 1).reshape(c, -1) noise = noise.reshape(c, b, h, w).swapaxes(1, 0) - return noise.mul_(self.factor) + return scale_noise(noise, self.factor, normalized=normalized) return sampler diff --git a/py/sonar.py b/py/sonar.py index b183ec0..ac6cc85 100644 --- a/py/sonar.py +++ b/py/sonar.py @@ -83,6 +83,7 @@ class SonarBase: sigma_max, seed=seed, cpu=True, + normalized=True, ) self.noise_sampler = noise_sampler return noise_sampler @@ -103,6 +104,7 @@ class SonarBase: None, seed=self.extra_args.get("seed"), cpu=True, + normalized=True, ) self.history_d = ns(None, None) else: @@ -160,34 +162,42 @@ class SonarGuidanceMixin: if self.ref_latent.device != x.device: self.ref_latent = self.ref_latent.to(device=x.device) if self.guidance.guidance_type == GuidanceType.LINEAR: - return self.guidance_linear(x) + return self.guidance_linear(x, self.ref_latent, self.guidance.factor) if self.guidance.guidance_type == GuidanceType.EULER: - return self.guidance_euler(step_index, x, denoised) + sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1] + return self.guidance_euler( + sigma, + sigma_next, + x, + denoised, + self.ref_latent, + self.guidance.factor, + ) raise ValueError("Sonar: Guidance: Unknown guidance type") + @staticmethod def guidance_euler( - self, - step_index: int, + sigma: Tensor, + sigma_next: Tensor, x: Tensor, denoised: Tensor, - ): + ref_latent: Tensor, + factor: float = 0.2, + ) -> Tensor: avg_t = denoised.mean(dim=[1, 2, 3], keepdim=True) std_t = denoised.std(dim=[1, 2, 3], keepdim=True) - ref_img_shift = self.ref_latent * std_t + avg_t - sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1] + ref_img_shift = ref_latent * std_t + avg_t d = sampling.to_d(x, sigma, ref_img_shift) - dt = (sigma_next - sigma) * self.guidance.factor + dt = (sigma_next - sigma) * factor return x + d * dt - def guidance_linear( - self, - x: Tensor, - ): + @staticmethod + def guidance_linear(x: Tensor, ref_latent: Tensor, factor: float = 0.2) -> Tensor: avg_t = x.mean(dim=[1, 2, 3], keepdim=True) std_t = x.std(dim=[1, 2, 3], keepdim=True) - ref_img_shift = self.ref_latent * std_t + avg_t - return (1.0 - self.guidance.factor) * x + self.guidance.factor * ref_img_shift + ref_img_shift = ref_latent * std_t + avg_t + return (1.0 - factor) * x + factor * ref_img_shift class SonarWithGuidance(SonarBase, SonarGuidanceMixin):