diff --git a/README.md b/README.md index a03eb1f..43a6031 100644 --- a/README.md +++ b/README.md @@ -164,7 +164,7 @@ Normal (non-sonar) Eular A. Not really a comparison with noise (think it would u ![Rainbow Intense](assets/example_images/noise/renoise_rainbow_intense.png) -#### Green_test_ +#### Green_test ![Green_test](assets/example_images/noise/renoise_green_test.png) diff --git a/py/noise.py b/py/noise.py index f6bd00b..7870895 100644 --- a/py/noise.py +++ b/py/noise.py @@ -16,7 +16,9 @@ from torch import FloatTensor, Generator, Tensor # This likely isn't correct. def scale_noise(noise, factor=1.0): + return noise * factor mean, std = noise.mean(), noise.std() + # print(f"factor={factor:.3}, mean={mean:.3}, std={std:.3}") noise = noise - mean if std >= 0.98: noise /= std @@ -115,7 +117,9 @@ class CustomNoiseChain: op.add, (ns(sigma, sigma_next) for ns in noise_samplers), ) - return scale_noise(result, scale) + result = scale_noise(result, scale) + # print("SCALING", scale, "sum", result.sum().item()) + return result return noise_sampler @@ -465,7 +469,9 @@ class NoiseSampler: args = ( self.transform(torch.as_tensor(s)) if s is not None else s for s in args ) - result = scale_noise(self.noise_sampler(*args, **kwargs), self.factor) + # print(">", self.factor) + result = self.noise_sampler(*args, **kwargs) * self.factor + # result = scale_noise(self.noise_sampler(*args, **kwargs), self.factor) if hasattr(result, "to"): return result.to(dtype=self.dtype, device=self.device) return result diff --git a/py/sonar.py b/py/sonar.py index da3797f..e2ea1c1 100644 --- a/py/sonar.py +++ b/py/sonar.py @@ -3,7 +3,7 @@ from __future__ import annotations from enum import Enum, auto -from typing import Any, NamedTuple +from typing import Any, Callable, NamedTuple import torch from comfy.k_diffusion import sampling @@ -44,14 +44,46 @@ class SonarConfig(NamedTuple): class SonarBase: - def __init__( - self, - cfg: SonarConfig, - ) -> None: + def __init__(self, cfg: SonarConfig) -> None: self.history_d = None self.cfg = cfg self.noise_sampler = None + def set_noise_sampler( + self, + x: Tensor, + sigmas, + noise_sampler: Callable | None, + seed: int | None = None, + ): + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + if noise_sampler is not None and self.cfg.noise_type not in ( + None, + noise.NoiseType.GAUSSIAN, + ): + # Possibly we should just use the supplied already-created noise sampler here. + raise ValueError( + "Unexpected noise_sampler presence with non-default noise type requested", + ) + if self.cfg.custom_noise: + noise_sampler = self.cfg.custom_noise.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=seed, + ) + elif noise_sampler is None and self.cfg.noise_type: + noise_sampler = noise.get_noise_sampler( + self.cfg.noise_type, + x, + sigma_min, + sigma_max, + seed=seed, + cpu=True, + ) + self.noise_sampler = noise_sampler + return noise_sampler + def init_hist_d(self, x: Tensor) -> None: if self.history_d is not None: return @@ -201,7 +233,7 @@ class SonarEuler(SonarSampler): ): self.init_hist_d(sample) - sigma = self.sigmas[step_index] + sigma, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1] gamma = ( min(self.s_churn / (len(self.sigmas) - 1), 2**0.5 - 1) @@ -212,8 +244,11 @@ class SonarEuler(SonarSampler): sigma_hat = sigma * (gamma + 1) if gamma > 0: - noise = torch.randn_like(sample) - + noise = ( + self.noise_sampler(sigma, sigma_to) + if self.noise_sampler + else torch.randn_like(sample) + ) eps = noise * self.s_noise sample = sample + eps * (sigma_hat**2 - sigma**2) ** 0.5 @@ -243,6 +278,7 @@ class SonarEuler(SonarSampler): extra_args=None, callback=None, disable=None, + noise_sampler: Callable | None = None, sonar_config=None, s_churn=0.0, s_tmin=0.0, @@ -263,6 +299,13 @@ class SonarEuler(SonarSampler): {} if extra_args is None else extra_args, sonar_config, ) + sonar.set_noise_sampler( + x, + sigmas, + noise_sampler, + seed=extra_args.get("seed"), + ) + # print("noise sampler:", noise_sampler, sonar.noise_sampler) for i in trange(len(sigmas) - 1, disable=disable): x, sigma, sigma_hat, denoised = sonar.step( @@ -285,14 +328,12 @@ class SonarEuler(SonarSampler): class SonarEulerAncestral(SonarSampler): def __init__( self, - noise_sampler, eta: float = 1.0, s_noise: float = 1.0, *args: list[Any], **kwargs: dict[str, Any], ): super().__init__(*args, **kwargs) - self.noise_sampler = noise_sampler self.eta = eta self.s_noise = s_noise @@ -342,7 +383,7 @@ class SonarEulerAncestral(SonarSampler): sonar_config=None, eta=1.0, s_noise=1.0, - noise_sampler=None, + noise_sampler: Callable | None = None, ): if sonar_config is None: sonar_config = SonarConfig() @@ -354,27 +395,8 @@ class SonarEulerAncestral(SonarSampler): raise ValueError( "Unexpected noise_sampler presence with non-default noise type requested", ) - sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() - seed = extra_args.get("seed") - if sonar_config.custom_noise: - noise_sampler = sonar_config.custom_noise.make_noise_sampler( - x, - sigma_min, - sigma_max, - seed=seed, - ) - elif noise_sampler is None: - noise_sampler = noise.get_noise_sampler( - sonar_config.noise_type, - x, - sigma_min, - sigma_max, - seed=seed, - cpu=True, - ) s_in = x.new_ones([x.shape[0]]) sonar = cls( - noise_sampler, eta, s_noise, model, @@ -383,6 +405,12 @@ class SonarEulerAncestral(SonarSampler): {} if extra_args is None else extra_args, sonar_config, ) + sonar.set_noise_sampler( + x, + sigmas, + noise_sampler, + seed=extra_args.get("seed"), + ) for i in trange(len(sigmas) - 1, disable=disable): x, sigma, sigma_hat, denoised = sonar.step( @@ -405,14 +433,12 @@ class SonarEulerAncestral(SonarSampler): class SonarDPMPPSDE(SonarSampler): def __init__( self, - noise_sampler, eta: float = 1.0, s_noise: float = 1.0, *args: list[Any], **kwargs: dict[str, Any], ): super().__init__(*args, **kwargs) - self.noise_sampler = noise_sampler self.eta = eta self.s_noise = s_noise @@ -543,27 +569,9 @@ class SonarDPMPPSDE(SonarSampler): raise ValueError( "Unexpected noise_sampler presence with non-default noise type requested", ) - sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() - seed = extra_args.get("seed") - if sonar_config.custom_noise: - noise_sampler = sonar_config.custom_noise.make_noise_sampler( - x, - sigma_min, - sigma_max, - seed=seed, - ) - elif noise_sampler is None: - noise_sampler = noise.get_noise_sampler( - sonar_config.noise_type, - x, - sigma_min, - sigma_max, - seed=seed, - cpu=True, - ) + s_in = x.new_ones([x.shape[0]]) sonar = cls( - noise_sampler, eta, s_noise, model, @@ -572,6 +580,12 @@ class SonarDPMPPSDE(SonarSampler): {} if extra_args is None else extra_args, sonar_config, ) + sonar.set_noise_sampler( + x, + sigmas, + noise_sampler, + seed=extra_args.get("seed"), + ) for i in trange(len(sigmas) - 1, disable=disable): x, sigma, sigma_hat, denoised = sonar.step(