From 227be2c85a8209335c12ab6de74e19ee7f192b8a Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 24 Feb 2024 09:59:45 -0700 Subject: [PATCH] Refactor noise: stage 2 --- __init__.py | 1 + changelog.md | 5 ++ py/nodes.py | 152 ++++++++++++++++++++++++++++++++++++++++++++++++++- py/noise.py | 21 ++++--- py/sonar.py | 26 ++++++--- 5 files changed, 190 insertions(+), 15 deletions(-) diff --git a/__init__.py b/__init__.py index 4a412fb..e17172d 100644 --- a/__init__.py +++ b/__init__.py @@ -6,6 +6,7 @@ NODE_CLASS_MAPPINGS = { "SamplerSonarEuler": nodes.SamplerNodeSonarEuler, "SamplerSonarEulerA": nodes.SamplerNodeSonarEulerAncestral, "SamplerSonarDPMPPSDE": nodes.SamplerNodeSonarDPMPPSDE, + "SamplerConfigOverride": nodes.SamplerNodeConfigOverride, "NoisyLatentLike": nodes.NoisyLatentLikeNode, "SonarCustomNoise": nodes.SonarCustomNoiseNode, "SonarGuidanceConfig": nodes.GuidanceConfigNode, diff --git a/changelog.md b/changelog.md index ee12cef..b7c435e 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,11 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20240224 + +* Refactored noise generation functions (will break seeds). +* Added `SamplerOverride` node. + ## 20240210 * Added `SonarCustomNoise` node. diff --git a/py/nodes.py b/py/nodes.py index fa98c59..9bee646 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -1,5 +1,8 @@ from __future__ import annotations +import inspect +from typing import Any, Callable + import torch from comfy import samplers @@ -55,7 +58,7 @@ class NoisyLatentLikeNode: latent["samples"], None, None, - seed=None, + seed=seed, cpu=True, ) randst = torch.random.get_rng_state() @@ -393,3 +396,150 @@ class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEuler): }, ), ) + + +class SamplerNodeConfigOverride: + KWARG_OVERRIDES = ("s_noise", "eta", "s_churn", "r", "solver_type") + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sampler": ("SAMPLER",), + "eta": ( + "FLOAT", + { + "default": 1.0, + "step": 0.01, + "round": False, + }, + ), + "s_noise": ( + "FLOAT", + { + "default": 1.0, + "step": 0.01, + "round": False, + }, + ), + "s_churn": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "step": 0.01, + "round": False, + }, + ), + "r": ( + "FLOAT", + { + "default": 0.5, + "step": 0.01, + "round": False, + }, + ), + "sde_solver": (("midpoint", "heun"),), + }, + "optional": { + "noise_type": (tuple(t.name.lower() for t in noise.NoiseType),), + "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), + }, + } + + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler( + self, + sampler, + eta, + s_noise, + s_churn, + r, + sde_solver, + noise_type=None, + custom_noise_opt=None, + ): + return ( + samplers.KSAMPLER( + self.sampler_function, + extra_options=sampler.extra_options + | { + "override_sampler_cfg": { + "sampler": sampler, + "noise_type": noise.NoiseType[noise_type.upper()] + if noise_type is not None + else None, + "custom_noise": custom_noise_opt, + "s_noise": s_noise, + "eta": eta, + "s_churn": s_churn, + "r": r, + "solver_type": sde_solver, + }, + }, + inpaint_options=sampler.inpaint_options | {}, + ), + ) + + @classmethod + @torch.no_grad() + def sampler_function( + cls, + model, + x, + sigmas, + *args: list[Any], + override_sampler_cfg: dict[str, Any] | None = None, + noise_sampler: Callable | None = None, + extra_args: dict[str, Any] | None = None, + **kwargs: dict[str, Any], + ): + if not override_sampler_cfg: + raise ValueError("Override sampler config missing!") + if extra_args is None: + extra_args = {} + cfg = override_sampler_cfg + sampler, noise_type, custom_noise = ( + cfg["sampler"], + cfg.get("noise_type"), + cfg.get("custom_noise"), + ) + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + seed = extra_args.get("seed") + if custom_noise is not None: + noise_sampler = custom_noise.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=seed, + ) + elif noise_type is not None: + noise_sampler = noise.get_noise_sampler( + noise_type, + x, + sigma_min, + sigma_max, + seed=seed, + cpu=True, + ) + sig = inspect.signature(sampler.sampler_function) + params = sig.parameters + kwargs = kwargs | {} + if "noise_sampler" in params: + kwargs["noise_sampler"] = noise_sampler + for k in cls.KWARG_OVERRIDES: + if k not in params or cfg.get(k) is None: + continue + kwargs[k] = cfg[k] + return sampler.sampler_function( + model, + x, + sigmas, + *args, + extra_args=extra_args, + **kwargs, + ) diff --git a/py/noise.py b/py/noise.py index d462f42..f6bd00b 100644 --- a/py/noise.py +++ b/py/noise.py @@ -14,6 +14,18 @@ from torch import FloatTensor, Generator, Tensor # ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311 +# This likely isn't correct. +def scale_noise(noise, factor=1.0): + mean, std = noise.mean(), noise.std() + noise = noise - mean + if std >= 0.98: + noise /= std + elif std < 0.98: + noise *= 1.0 + std + noise *= factor + return noise + + class NoiseType(Enum): GAUSSIAN = auto() UNIFORM = auto() @@ -103,11 +115,7 @@ class CustomNoiseChain: op.add, (ns(sigma, sigma_next) for ns in noise_samplers), ) - std = result.std() - if std > 1: - result /= std - result *= scale - return result + return scale_noise(result, scale) return noise_sampler @@ -457,8 +465,7 @@ class NoiseSampler: args = ( self.transform(torch.as_tensor(s)) if s is not None else s for s in args ) - result = self.noise_sampler(*args, **kwargs) - result *= 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 0cea765..da3797f 100644 --- a/py/sonar.py +++ b/py/sonar.py @@ -212,7 +212,7 @@ class SonarEuler(SonarSampler): sigma_hat = sigma * (gamma + 1) if gamma > 0: - noise = torch.randn_like(sample.shape) + noise = torch.randn_like(sample) eps = noise * self.s_noise sample = sample + eps * (sigma_hat**2 - sigma**2) ** 0.5 @@ -355,15 +355,21 @@ class SonarEulerAncestral(SonarSampler): "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) - else: + 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=extra_args.get("seed"), + seed=seed, cpu=True, ) s_in = x.new_ones([x.shape[0]]) @@ -538,15 +544,21 @@ class SonarDPMPPSDE(SonarSampler): "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) - else: + 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=extra_args.get("seed"), + seed=seed, cpu=True, ) s_in = x.new_ones([x.shape[0]])