Make custom noise more extensible (internal change)

This commit is contained in:
blepping
2024-03-02 04:48:54 -07:00
parent a2fd7118b5
commit 090df2280d
2 changed files with 86 additions and 30 deletions
+33 -10
View File
@@ -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):
+53 -20
View File
@@ -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
)