Author SHA1 Message Date
blepping a4fed311a8 Partial 5D latent support for custom noise types 2025-02-27 12:33:35 -07:00
blepping 2988afa34a Fix Bleh and Restart integration 2025-02-17 03:09:46 -07:00
blepping 68fc7418d1 Yet another noise generation refactor (#12)
* Refactor noise generation
Try to make option passing and CPU/GPU noise selection work
Add advanced custom noise node that allows for parameter passing
Add wavelet noise type

* Add WaveletFilteredNoise node, other fixes

* Fix Brownian arg passing

* Generalized distribution noise for most torch.distributions

* Distro noise improvements, add SonarAdvancedDistroNoise node

* More distributions!

* Add SonarResizedNoise node

* Momentum sampler refactor/improvements (I hope)

* Better approach to integration with external nodes
Documentation updates
Other cleanups

* Internal cleanups and refactoring.
Some integration improvements.
Bump date in changelog

* Add round and step to node FLOAT inputs that did not have it
2025-01-31 17:33:04 -07:00
5 changed files with 147 additions and 98 deletions
+4
View File
@@ -2,6 +2,10 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20250227
* Add 5D latent (video models) support for most custom noise types.
## 20250130
*Note*: May change seeds.
+23 -12
View File
@@ -2216,6 +2216,12 @@ class SonarBlendFilterNoiseNode(
@classmethod
def INPUT_TYPES(cls):
result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
bleh_filter_presets = (
() if bleh is None else tuple(bleh.py.latent_utils.FILTER_PRESETS.keys())
)
bleh_enhance_methods = (
() if bleh is None else ("none", *bleh.py.latent_utils.ENHANCE_METHODS)
)
result["required"] |= {
"sonar_custom_noise": (
WILDCARD_NOISE,
@@ -2224,7 +2230,7 @@ class SonarBlendFilterNoiseNode(
},
),
"blend_mode": (("simple_add", *utils.BLENDING_MODES.keys()),),
"ffilter": (tuple(bleh.py.latent_utils.FILTER_PRESETS.keys()),),
"ffilter": (bleh_filter_presets,),
"ffilter_custom": ("STRING", {"default": ""}),
"ffilter_scale": (
"FLOAT",
@@ -2250,7 +2256,7 @@ class SonarBlendFilterNoiseNode(
"INT",
{"default": 1, "min": 1, "max": 32},
),
"enhance_mode": (("none", *bleh.py.latent_utils.ENHANCE_METHODS),),
"enhance_mode": (bleh_enhance_methods,),
"enhance_strength": (
"FLOAT",
{
@@ -2385,11 +2391,18 @@ class KRestartSamplerCustomNoise(metaclass=IntegratedNode):
@classmethod
def INPUT_TYPES(cls):
get_normal_schedulers = getattr(
restart.nodes,
"get_supported_normal_schedulers",
restart.nodes.get_supported_restart_schedulers,
)
if restart is not None:
get_normal_schedulers = getattr(
restart.nodes,
"get_supported_normal_schedulers",
restart.nodes.get_supported_restart_schedulers,
)
restart_normal_schedulers = get_normal_schedulers()
restart_schedulers = restart.nodes.get_supported_restart_schedulers()
restart_default_segments = restart.restart_sampling.DEFAULT_SEGMENTS
else:
restart_default_segments = ""
restart_normal_schedulers = restart_schedulers = ()
return {
"required": {
"model": ("MODEL",),
@@ -2410,7 +2423,7 @@ class KRestartSamplerCustomNoise(metaclass=IntegratedNode):
},
),
"sampler": ("SAMPLER",),
"scheduler": (get_normal_schedulers(),),
"scheduler": (restart_normal_schedulers,),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"latent_image": ("LATENT",),
@@ -2420,13 +2433,11 @@ class KRestartSamplerCustomNoise(metaclass=IntegratedNode):
"segments": (
"STRING",
{
"default": restart.restart_sampling.DEFAULT_SEGMENTS,
"default": restart_default_segments,
"multiline": False,
},
),
"restart_scheduler": (
restart.nodes.get_supported_restart_schedulers(),
),
"restart_scheduler": (restart_schedulers,),
"chunked_mode": ("BOOLEAN", {"default": True}),
},
"optional": {
+26 -24
View File
@@ -1388,8 +1388,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="onef_pinkishgreenish",
noise_mix=(
partial(OneFNoiseGenerator, alpha=0.5),
partial(OneFNoiseGenerator, alpha=-0.5),
(OneFNoiseGenerator, {"alpha": 0.5}, None),
(OneFNoiseGenerator, {"alpha": -0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
@@ -1399,8 +1399,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="onef_pinkish_mix",
noise_mix=(
(partial(OneFNoiseGenerator, alpha=0.5), lambda t: t.mul_(-1.0)),
partial(OneFNoiseGenerator, alpha=0.5),
(OneFNoiseGenerator, {"alpha": 0.5}, lambda t: t.mul_(-1.0)),
(OneFNoiseGenerator, {"alpha": 0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
@@ -1410,8 +1410,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="onef_greenish_mix",
noise_mix=(
(partial(OneFNoiseGenerator, alpha=0.5), lambda t: t.mul_(-1.0)),
partial(OneFNoiseGenerator, alpha=0.5),
(OneFNoiseGenerator, {"alpha": 0.5}, lambda t: t.mul_(-1.0)),
(OneFNoiseGenerator, {"alpha": 0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
@@ -1455,8 +1455,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="rainbow_mild",
noise_mix=(
(GreenTestNoiseGenerator, lambda t: t.mul_(0.55)),
(GreenTestNoiseGenerator, lambda t: t.mul_(0.7)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.55)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.7)),
),
output_fun=lambda t: t.mul_(1.15),
),
@@ -1466,8 +1466,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="rainbow_intense",
noise_mix=(
(GreenTestNoiseGenerator, lambda t: t.mul_(0.75)),
(GreenTestNoiseGenerator, lambda t: t.mul_(0.5)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.75)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.5)),
),
output_fun=lambda t: t.mul_(1.15),
),
@@ -1502,8 +1502,8 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
MixedNoiseGenerator,
name="pyramid_mix",
noise_mix=(
(partial(PyramidNoiseGenerator, discount=0.6), lambda t: t.mul_(0.2)),
(partial(PyramidNoiseGenerator, discount=0.6), lambda t: t.mul_(-0.8)),
(PyramidNoiseGenerator, {"discount": 0.6}, lambda t: t.mul_(0.2)),
(PyramidNoiseGenerator, {"discount": 0.6}, lambda t: t.mul_(-0.8)),
),
),
),
@@ -1513,11 +1513,13 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
name="pyramid_mix_area",
noise_mix=(
(
partial(PyramidNoiseGenerator, discount=0.5, upscale_mode="area"),
PyramidNoiseGenerator,
{"discount": 0.5, "upscale_mode": "area"},
lambda t: t.mul_(0.2),
),
(
partial(PyramidNoiseGenerator, discount=0.5, upscale_mode="area"),
PyramidNoiseGenerator,
{"discount": 0.5, "upscale_mode": "area"},
lambda t: t.mul_(-0.8),
),
),
@@ -1529,19 +1531,19 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = {
name="pyramid_mix_bislerp",
noise_mix=(
(
partial(
PyramidNoiseGenerator,
discount=0.5,
upscale_mode="bislerp",
),
PyramidNoiseGenerator,
{
"discount": 0.5,
"upscale_mode": "bislerp",
},
lambda t: t.mul_(0.2),
),
(
partial(
PyramidNoiseGenerator,
discount=0.5,
upscale_mode="bislerp",
),
PyramidNoiseGenerator,
{
"discount": 0.5,
"upscale_mode": "bislerp",
},
lambda t: t.mul_(-0.8),
),
),
+93 -61
View File
@@ -79,7 +79,6 @@ class NoiseError(Exception):
class NoiseGenerator:
name = "unknown"
SAVE_X = False
MIN_DIMS = 1
MAX_DIMS = 0
@@ -100,7 +99,7 @@ class NoiseGenerator:
setattr(self, k, kwarg_params.pop(k))
self.options = kwarg_params
self.update_x(x)
print("CREATE NG", self, kwargs)
# print("CREATE NG", self, kwargs)
@classmethod
@property
@@ -114,11 +113,13 @@ class NoiseGenerator:
}
def update_x(self, x):
self.x = None if not self.SAVE_X else x.detach().clone()
self.shape = x.shape
self.batch, self.channels, self.height, self.width = (
x.shape if x.ndim == 4 else (None,) * 4
)
if x.ndim in {4, 5}:
self.batch, self.channels = x.shape[:2]
self.height, self.width = x.shape[-2:]
self.frames = x.shape[-3] if x.ndim == 5 else None
else:
self.batch = self.channels = self.frames = self.height = self.width = None
self.device = x.device
self.gen_device = torch.device("cpu") if self.cpu else self.device
self.layout = x.layout
@@ -133,14 +134,11 @@ class NoiseGenerator:
layout=self.layout,
device=self.gen_device,
)
# print("GEN NOISE", noise)
if to_device and noise.device != self.device:
noise = tensor_to(noise, self.device)
# print("MADE NOISE", noise)
return noise
def output_hook(self, noise):
# print("NOISE OUT1", noise)
if noise.device != self.device:
noise = tensor_to(noise, self.device)
return scale_noise(
@@ -165,9 +163,35 @@ class NoiseGenerator:
return f"<NoiseGenerator({self.name}): device={self.device}, shape={self.shape}, dtype={self.dtype}, {pretty_params}>"
class MixedNoiseGenerator(NoiseGenerator):
MIN_DIMS = MAX_DIMS = 4
class FramesToChannelsNoiseGenerator(NoiseGenerator):
MIN_DIMS = 4
MAX_DIMS = 5
def get_adjusted_shape(self):
if self.frames:
return (self.batch, self.channels * self.frames, self.height, self.width)
return (self.batch, self.channels, self.height, self.width)
def fix_output_frames(self, noise):
if not self.frames:
return noise
return noise.reshape(
self.batch,
self.channels,
self.frames,
self.height,
self.width,
)
def rand_like(self, *args, **kwargs):
noise = super().rand_like(*args, **kwargs)
adjusted_shape = self.get_adjusted_shape()
if noise.shape != adjusted_shape:
return noise.reshape(*adjusted_shape)
return noise
class MixedNoiseGenerator(NoiseGenerator):
@classmethod
@property
def ng_params(cls):
@@ -180,15 +204,20 @@ class MixedNoiseGenerator(NoiseGenerator):
}
def __init__(self, x, *args, **kwargs):
min_dim = max_dim = None
self.name = kwargs["name"]
for item in kwargs["noise_mix"]:
ng_class = item[0] if isinstance(item, (tuple, list)) else item
cmin, cmax = ng_class.MIN_DIMS, ng_class.MAX_DIMS
min_dim = max(min_dim if min_dim is not None else cmin, cmin)
max_dim = min(max_dim if max_dim is not None else cmax, cmax)
self.MIN_DIMS = min_dim
self.MAX_DIMS = max_dim
super().__init__(x, *args, **kwargs)
ng_list = []
for item in self.noise_mix:
if isinstance(item, (tuple, list)):
ng_class, transform_fun = item
else:
ng_class, transform_fun = item, None
for ng_class, ng_class_kwargs, transform_fun in self.noise_mix:
ng_kwargs = {k: v for k, v in kwargs.items() if k in self.pass_args}
ng_list.append((ng_class(x, **ng_kwargs), transform_fun))
ng_list.append((ng_class(x, **ng_class_kwargs, **ng_kwargs), transform_fun))
self.ng_list = ng_list
def generate(self, *args):
@@ -242,9 +271,8 @@ class BrownianNoiseGenerator(NoiseGenerator):
return self.brownian_tree_ns(*args)
class PerlinOldNoiseGenerator(NoiseGenerator):
class PerlinOldNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "perlin_old"
MIN_DIMS = MAX_DIMS = 4
@classmethod
@property
@@ -437,18 +465,18 @@ class PerlinOldNoiseGenerator(NoiseGenerator):
blend = utils.BLENDING_MODES[self.blend_mode]
noise = self.rand_like(fun=torch.rand).div_(self.div_fac)
_batch, channels, noise_height, noise_width = noise.shape
channels, height, width = noise.shape[1:]
for _ in range(self.iterations):
noise += self.perlin_noise(
(noise_height, noise_width),
(noise_height, noise_width),
batch_size=channels, # This should be the number of channels.
(height, self.width),
(height, width),
batch_size=channels,
blend=blend,
dtype=noise.dtype,
layout=noise.layout,
device=noise.device,
)
return noise
return self.fix_output_frames(noise)
class UniformNoiseGenerator(NoiseGenerator):
@@ -473,9 +501,8 @@ class UniformNoiseGenerator(NoiseGenerator):
)
class HighresPyramidNoiseGenerator(NoiseGenerator):
class HighresPyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "highres_pyramid"
MIN_DIMS = MAX_DIMS = 4
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
@@ -502,15 +529,10 @@ class HighresPyramidNoiseGenerator(NoiseGenerator):
}
def generate(self, s, sn):
(
b,
c,
h,
w,
) = self.shape # EDIT: w and h get over-written, rename for a different variant!
adjusted_shape = self.get_adjusted_shape()
b, c, h, w = adjusted_shape
orig_w, orig_h = w, h
noise = self.uniform_ng(s, sn)
noise = self.uniform_ng(s, sn).reshape(*adjusted_shape)
rs = (
torch.rand(
self.iterations,
@@ -531,12 +553,11 @@ class HighresPyramidNoiseGenerator(NoiseGenerator):
).mul_(self.discount**i)
if h >= orig_h * 15 or w >= orig_w * 15:
break # Lowest resolution is 1x1
return noise
return self.fix_output_frames(noise)
class PyramidOldNoiseGenerator(NoiseGenerator):
class PyramidOldNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "pyramid_old"
MIN_DIMS = MAX_DIMS = 4
@classmethod
@property
@@ -549,10 +570,11 @@ class PyramidOldNoiseGenerator(NoiseGenerator):
}
def generate(self, *_args):
b, c, h, w = self.shape
adjusted_shape = self.get_adjusted_shape()
b, c, h, w = adjusted_shape
orig_h, orig_w = h, w
noise = torch.zeros(
size=self.shape,
size=adjusted_shape,
dtype=self.dtype,
layout=self.layout,
device=self.gen_device,
@@ -574,12 +596,11 @@ class PyramidOldNoiseGenerator(NoiseGenerator):
orig_h,
mode=self.upscale_mode,
).mul_(self.discount**i)
return noise
return self.fix_output_frames(noise)
class PyramidNoiseGenerator(NoiseGenerator):
class PyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "pyramid"
MIN_DIMS = MAX_DIMS = 4
@classmethod
@property
@@ -592,12 +613,10 @@ class PyramidNoiseGenerator(NoiseGenerator):
# Modified from https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2
def generate(self, *_args):
b, c, w, h = (
self.shape
) # NOTE: w and h get over-written, rename for a different variant!
orig_w, orig_h = w, h
noise = self.rand_like()
b, c, h, w = noise.shape
orig_w, orig_h = w, h
for i in range(self.iterations):
r = (
torch.rand(1, generator=self.generator).cpu().item() * 2 + 2
@@ -621,7 +640,7 @@ class PyramidNoiseGenerator(NoiseGenerator):
)
if w == 1 or h == 1:
break # Lowest resolution is 1x1
return noise
return self.fix_output_frames(noise)
class StudentTNoiseGenerator(NoiseGenerator):
@@ -653,9 +672,10 @@ class StudentTNoiseGenerator(NoiseGenerator):
return torch.copysign(torch.pow(torch.abs(noise), self.pow_fac), noise)
class GreenTestNoiseGenerator(NoiseGenerator):
class GreenTestNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "green_test"
MIN_DIMS = MAX_DIMS = 4
MIN_DIMS = 4
MAX_DIMS = 5
@classmethod
@property
@@ -677,7 +697,7 @@ class GreenTestNoiseGenerator(NoiseGenerator):
power[0, 0] = self.power_base
noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power))
noise *= scale / noise.std()
return torch.real(noise)
return self.fix_output_frames(torch.real(noise))
class PinkOldNoiseGenerator(NoiseGenerator):
@@ -694,9 +714,10 @@ class PinkOldNoiseGenerator(NoiseGenerator):
return self.rand_like() * spectral_density
class OneFNoiseGenerator(NoiseGenerator):
class OneFNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "onef"
MIN_DIMS = MAX_DIMS = 4
MIN_DIMS = 4
MAX_DIMS = 5
@classmethod
@property
@@ -712,19 +733,19 @@ class OneFNoiseGenerator(NoiseGenerator):
# Referenced from: https://github.com/WASasquatch/PowerNoiseSuite
def generate(self, *_args):
batch, _channels, height, width = self.shape
# batch, _channels, height, width = self.shape
noise = self.rand_like()
freq_x = tensor_to(torch.fft.fftfreq(height, self.hfac), noise)
freq_y = tensor_to(torch.fft.fftfreq(width, self.wfac), noise)
freq_x = tensor_to(torch.fft.fftfreq(self.height, self.hfac), noise)
freq_y = tensor_to(torch.fft.fftfreq(self.width, self.wfac), noise)
fx, fy = torch.meshgrid(freq_x, freq_y, indexing="ij")
power = (fx**2 + fy**2) ** (-self.alpha / 2.0)
if self.k != 0:
power = self.k / power
power[0, 0] = self.base_power
power = power.unsqueeze(0).expand(batch, 1, height, width)
power = power.unsqueeze(0).expand(self.batch, 1, self.height, self.width)
noise_fft = torch.fft.fftn(noise)
noise_fft /= (
@@ -733,7 +754,7 @@ class OneFNoiseGenerator(NoiseGenerator):
else power.to(noise_fft.dtype)
)
return torch.fft.ifftn(noise_fft).real
return self.fix_output_frames(torch.fft.ifftn(noise_fft).real)
class PowerLawNoiseGenerator(NoiseGenerator):
@@ -1258,9 +1279,10 @@ class PowerOldNoiseGenerator(NoiseGenerator):
# Idea from https://github.com/ClownsharkBatwing/RES4LYF/ (wave and mode defaults also from that source)
class WaveletNoiseGenerator(NoiseGenerator):
class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "wavelet"
MIN_DIMS = MAX_DIMS = 4
MIN_DIMS = 4
MAX_DIMS = 5
def __init__(self, *args, **kwargs):
if not HAVE_WAVELETS:
@@ -1311,11 +1333,21 @@ class WaveletNoiseGenerator(NoiseGenerator):
}
def generate(self, *args):
adjusted_shape = self.get_adjusted_shape()
noise = (
self.rand_like()
if self.noise_sampler is None
else self.noise_sampler(*args)
)
if noise.shape != adjusted_shape:
noise = noise.reshape(*adjusted_shape)
if self.frames:
noise = noise.reshape(
self.batch,
self.channels * self.frames,
self.height,
self.width,
)
yl, yh = self.wavelet_forward(noise)
if self.yl_scale != 1:
yl *= self.yl_scale
@@ -1332,7 +1364,7 @@ class WaveletNoiseGenerator(NoiseGenerator):
for lidx in range(min(ht.shape[2], len(hscale))):
# print(">> SCALE IDX", lidx)
ht[:, :, lidx, :, :] *= hscale[lidx]
return self.wavelet_inverse((yl, yh))
return self.fix_output_frames(self.wavelet_inverse((yl, yh)))
__all__ = (
+1 -1
View File
@@ -636,7 +636,7 @@ class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
"max": 100.0,
"step": 0.001,
"round": False,
"tooltip": "Attempts to desaturate thelatent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
"tooltip": "Attempts to desaturate the latent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
},
),
"channel_correlation": (