From 090df2280d225b9901226741db0890817efbb4df Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 2 Mar 2024 04:48:54 -0700 Subject: [PATCH] Make custom noise more extensible (internal change) --- py/nodes.py | 43 +++++++++++++++++++++++-------- py/noise.py | 73 ++++++++++++++++++++++++++++++++++++++--------------- 2 files changed, 86 insertions(+), 30 deletions(-) diff --git a/py/nodes.py b/py/nodes.py index 9bee646..c554167 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -1,5 +1,6 @@ from __future__ import annotations +import abc import inspect from typing import Any, Callable @@ -70,7 +71,11 @@ class NoisyLatentLikeNode: return ({"samples": result},) -class SonarCustomNoiseNode: +class SonarCustomNoiseNodeBase(abc.ABC): + @abc.abstractmethod + def get_item_class(self): + raise NotImplementedError + @classmethod def INPUT_TYPES(cls): return { @@ -95,13 +100,6 @@ class SonarCustomNoiseNode: "round": False, }, ), - "noise_type": ( - tuple( - t.name.lower() - for t in noise.NoiseType - if t is not noise.NoiseType.BROWNIAN - ), - ), }, "optional": { "sonar_custom_noise_opt": ("SONAR_CUSTOM_NOISE",), @@ -112,17 +110,42 @@ class SonarCustomNoiseNode: CATEGORY = "advanced/noise" FUNCTION = "go" - def go(self, factor, rescale, noise_type, sonar_custom_noise_opt=None): + def go( + self, + factor, + rescale, + sonar_custom_noise_opt=None, + **kwargs: dict[str, Any], + ): nis = ( sonar_custom_noise_opt.clone() if sonar_custom_noise_opt else noise.CustomNoiseChain() ) if factor != 0: - nis.add(noise.CustomNoiseItem(factor, noise_type)) + nis.add(self.get_item_class()(factor, **kwargs)) return (nis if rescale == 0 else nis.rescaled(rescale),) +class SonarCustomNoiseNode(SonarCustomNoiseNodeBase): + @classmethod + def INPUT_TYPES(cls): + result = super().INPUT_TYPES() + result["required"] |= { + "noise_type": ( + tuple( + t.name.lower() + for t in noise.NoiseType + if t is not noise.NoiseType.BROWNIAN + ), + ), + } + return result + + def get_item_class(self): + return noise.CustomNoiseItem + + class GuidanceConfigNode: @classmethod def INPUT_TYPES(cls): diff --git a/py/noise.py b/py/noise.py index ff5fe2a..e80d6b2 100644 --- a/py/noise.py +++ b/py/noise.py @@ -1,6 +1,7 @@ # 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 @@ -16,7 +17,6 @@ from torch import FloatTensor, Generator, Tensor def scale_noise(noise, factor=1.0): mean, std = noise.mean(), noise.std() - # print(f"NOISE * {factor:.3}: std={std:.3}, mean={mean:.3}") return (noise - mean).div_(std).mul_(factor) @@ -42,10 +42,56 @@ class NoiseError(Exception): pass -class CustomNoiseItem: - def __init__(self, factor, noise_type): +class CustomNoiseItemBase(abc.ABC): + def __init__(self, factor, **kwargs): self.factor = factor - self.noise_type = noise_type + self.keys = set(kwargs.keys()) + for k, v in kwargs.items(): + setattr(self, k, v) + + def clone(self): + return self.__class__(self.factor, **{k: getattr(self, k) for k in self.keys}) + + def set_factor(self, factor): + self.factor = factor + return self + + @abc.abstractmethod + def make_noise_sampler( + self, + x: Tensor, + sigma_min=None, + sigma_max=None, + seed=None, + cpu=True, + ): + raise NotImplementedError + + +class CustomNoiseItem(CustomNoiseItemBase): + def __init__(self, factor, **kwargs): + super().__init__(factor, **kwargs) + if getattr(self, "noise_type", None) is None: + raise ValueError("Noise type required!") + + @torch.no_grad() + def make_noise_sampler( + self, + x: Tensor, + sigma_min=None, + sigma_max=None, + seed=None, + cpu=True, + ): + return get_noise_sampler( + self.noise_type, + x, + sigma_min, + sigma_max, + seed=seed, + cpu=cpu, + factor=self.factor, + ) class CustomNoiseChain: @@ -54,7 +100,7 @@ class CustomNoiseChain: def clone(self): return CustomNoiseChain( - [CustomNoiseItem(i.factor, i.noise_type) for i in self.items], + [i.clone() for i in self.items], ) def add(self, item): @@ -65,20 +111,9 @@ class CustomNoiseChain: divisor = total / scale divisor = divisor if divisor != 0 else 1.0 return CustomNoiseChain( - [CustomNoiseItem(i.factor / divisor, i.noise_type) for i in self.items], + [i.clone().set_factor(i.factor / divisor) for i in self.items], ) - def __call__( - self, - x: Tensor, - sigma_min=None, - sigma_max=None, - seed=None, - transform=lambda x: x, - cpu=True, - ): - pass - @torch.no_grad() def make_noise_sampler( self, @@ -89,14 +124,12 @@ class CustomNoiseChain: cpu=True, ) -> Callable: noise_samplers = tuple( - get_noise_sampler( - i.noise_type, + i.make_noise_sampler( x, sigma_min, sigma_max, seed=seed, cpu=cpu, - factor=i.factor, ) for i in self.items )