Noise refactor (#1)
* Refactor noise: stage 1 * Refactor noise: stage 2 * Refactor noise: stage 3 * Refactor noise: stage 4 * Update documentation and changelog
This commit is contained in:
@@ -4,6 +4,8 @@ A janky implementation of Sonar sampling (momentum-based sampling) for [ComfyUI]
|
||||
|
||||
Currently supports Euler, Euler Ancestral, and DPM++ SDE sampling.
|
||||
|
||||
See the [ChangeLog](changelog.md) for recent user-visible changes.
|
||||
|
||||
## Description
|
||||
|
||||
See https://github.com/Kahsolt/stable-diffusion-webui-sonar for a more in-depth explanation.
|
||||
@@ -22,11 +24,15 @@ You can also just choose `sonar_euler`, `sonar_euler_ancestral` or `sonar_dpmpp_
|
||||
|
||||
## Nodes
|
||||
|
||||
1. `SamplerSonarEuler` — Custom sampler node that combines Euler sampling and momentum and optionally guidance. A bit boring compared to the ancestral version but it has predictability going for it. You can possibly try setting init type to `RAND` and using different noise types, however this sampler seems _very_ sensitive to that init type. You may want to set direction to a very low value like `0.05` or `-0.15` when using the `RAND` init type.
|
||||
2. `SamplerSonarEulerAncestral` — Ancestral version of the above. Same features, just with ancestral Euler.
|
||||
4. `SonarGuidanceConfig` — You can optionally plug this into the Sonar sampler nodes. See the [Guidance](#guidance) section below.
|
||||
5. `NoisyLatentLike` — If you give it a latent (or latent batch) it'll return a noisy latent of the same shape. Allows specifying all the custom noise types except `brownian` which has some special requirements. Provided just because the noise generation functions are conveniently available. You can also use this as a reference latent with `SonarGuidanceConfig` node and depending on the strength it can act like variation seed (you'd change the seed in the `NoisyLatentLike` node). *Note*: The seed stuff may or may not work correctly.
|
||||
6. `SamplerSonarDPMPPSDE` — This one is extra experimental but it is an attempt to add moment and guidance to the DPM++ SDE sampler. It may not work correctly but you can sample stuff with it and get interesting results. I actually really like this one, and you can get away with more extreme stuff like `green_test` noise and still produce reasonable results. You may want to use the `BlehDiscardPenultimateSigma` node from my [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) collection if you find the result seems a bit washed out and b lurry.
|
||||
* `SamplerSonarEuler` — Custom sampler node that combines Euler sampling and momentum and optionally guidance. A bit boring compared to the ancestral version but it has predictability going for it. You can possibly try setting init type to `RAND` and using different noise types, however this sampler seems _very_ sensitive to that init type. You may want to set direction to a very low value like `0.05` or `-0.15` when using the `RAND` init type. Setting `momentum=1` is the same as disabling momentum, so this sampler with `momentum=1` is basically the same as the basic `euler` sampler.
|
||||
* `SamplerSonarEulerAncestral` — Ancestral version of the above. Same features, just with ancestral Euler.
|
||||
* `SonarGuidanceConfig` — You can optionally plug this into the Sonar sampler nodes. See the [Guidance](#guidance) section below.
|
||||
* `NoisyLatentLike` — If you give it a latent (or latent batch) it'll return a noisy latent of the same shape. Allows specifying all the custom noise types except `brownian` which has some special requirements. Provided just because the noise generation functions are conveniently available. You can also use this as a reference latent with `SonarGuidanceConfig` node and depending on the strength it can act like variation seed (you'd change the seed in the `NoisyLatentLike` node). *Note*: The seed stuff may or may not work correctly.
|
||||
* `SamplerSonarDPMPPSDE` — This one is extra experimental but it is an attempt to add moment and guidance to the DPM++ SDE sampler. It may not work correctly but you can sample stuff with it and get interesting results. I actually really like this one, and you can get away with more extreme stuff like `green_test` noise and still produce reasonable results. You may want to use the `BlehDiscardPenultimateSigma` node from my [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) collection if you find the result seems a bit washed out and blurry.
|
||||
* `SamplerConfigOverride` — can be used to override configuration settings for other samplers, including the noise type. For example, you could force `euler_ancestral` to use a different noise type. It's also possible to override other settings like `s_noise`, etc. *Note*: The wrapper inspects the sampling function's arguments to see what it supports, so you should connect the sampler directly to this rather than having other nodes (like a different sampler wrapper) in between.
|
||||
* `SonarCustomNoise` — See the [Noise](#noise) section below.
|
||||
|
||||
*Note*: `NoisyLatentLike` and `SamplerConfigOverride` are candidates for moving to a different project. They're just here at the moment because the noise generation functions are readily available.
|
||||
|
||||
## Parameters
|
||||
|
||||
@@ -51,19 +57,22 @@ I basically just copied a bunch of noise functions without really knowing what t
|
||||
3. `brownian`: This is the noise type SDE samplers use.
|
||||
4. `perlin`
|
||||
5. `studentt`: There's a comment that says it may enhance subject details. It seemed to produce a fairly dark result.
|
||||
6. `studentt_test`: An experiment that may be removed, it doesn't seem to be adding enough noise. You can possibly compensate by increasing `s_noise`.
|
||||
7. `pink`
|
||||
8. `highres_pyramid`: Not extensively tested, but it is slower than the other noise types. I would guess it does something like enhance details.
|
||||
9. `laplacian`
|
||||
10. `power`
|
||||
11. `rainbow_mild` and `rainbow_intense`: A combination of green (-ish, the implementation may be broken) noise plus perlin noise. Very colorful results.
|
||||
12. `green_test`: Even more rainbow-y than the rainbow noise types. It _probably_ isn't working correctly, but the results are very interesting and colorful. Depending on the model, it may not work well for an initial generation but may be worth trying with img2img type workflows.
|
||||
6. `pink`
|
||||
7. `highres_pyramid`: Not extensively tested, but it is slower than the other noise types. I would guess it does something like enhance details.
|
||||
8. `laplacian`
|
||||
9. `power`
|
||||
10. `rainbow_mild` and `rainbow_intense`: A combination of green (-ish, the implementation may be broken) noise plus perlin noise. Very colorful results.
|
||||
11. `green_test`: Even more rainbow-y than the rainbow noise types. It _probably_ isn't working correctly, but the results are very interesting and colorful. Depending on the model, it may not work well for an initial generation but may be worth trying with img2img type workflows.
|
||||
|
||||
You can scroll down to the the [Examples](#examples) section near the bottom to see some example generations with different noise types.
|
||||
|
||||
The sampler and `NoisyLatentLike` nodes now take an optional `SonarCustomNoise` input. You can chain `SonarCustomNoise` nodes together to mix different types of noise, similar to how some of the built in ones. It shouldn't matter what order the noise types are chained. If `rescale` is set to `0.0` no rescaling will occur. `factor` is the proportion of that type of noise you want. If you want to use `rescale` it should be on the node that you are plugging into a sampler. Just for example if you had two `SonarCustomNoise` nodes both with `factor=0.7` and `rescale=1.0` on the last one, it would be effectively the same as if you'd used `factor=0.5` and `rescale=1.0` doesn't actually do anything. You can also rescale to values above `1.0` — the result is more noise, similar to increasing `s_noise` above `1.0` on a sampler. The simple explanation is `rescale` means you don't have to make sure the `factor`s add up to the scale you want (which normally would be `1.0`).
|
||||
|
||||
**Note**: If you connect the optional `SonarCustomNoise` node to a Sonar sampler or the `NoisyLatentLike` node it will override the noise type selected in the node.
|
||||
**Note**: If you connect the optional `SonarCustomNoise` node to a Sonar sampler, the `NoisyLatentLike` node or the `SamplerConfigOverride` node, it will override the noise type selected in the node.
|
||||
|
||||
## Related
|
||||
|
||||
I also have some other ComfyUI nodes here: https://github.com/blepping/ComfyUI-bleh/
|
||||
|
||||
## Credits
|
||||
|
||||
@@ -141,11 +150,15 @@ Normal (non-sonar) Eular A. Not really a comparison with noise (think it would u
|
||||
|
||||
#### StudentT
|
||||
|
||||
**outdated**
|
||||
|
||||

|
||||
|
||||
|
||||
#### StudentT_test
|
||||
|
||||
**outdated**
|
||||
|
||||

|
||||
|
||||
#### Laplacian
|
||||
@@ -164,7 +177,7 @@ Normal (non-sonar) Eular A. Not really a comparison with noise (think it would u
|
||||
|
||||

|
||||
|
||||
#### Green_test_
|
||||
#### Green_test
|
||||
|
||||

|
||||
|
||||
@@ -203,10 +216,14 @@ These were generated with `s_noise=1.1` to make the noise effect more pronounced
|
||||
|
||||
#### StudentT
|
||||
|
||||
**outdated**
|
||||
|
||||

|
||||
|
||||
#### StudentT_test
|
||||
|
||||
**outdated**
|
||||
|
||||

|
||||
|
||||
#### Laplacian
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
|
||||
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
|
||||
|
||||
## 20240227
|
||||
|
||||
* Refactored noise generation functions (will break seeds).
|
||||
* Added `SamplerOverride` node.
|
||||
* `studentt` noise type replaced with `studentt_test` (the more correct version).
|
||||
|
||||
## 20240210
|
||||
|
||||
* Added `SonarCustomNoise` node.
|
||||
|
||||
+153
-3
@@ -1,5 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
from comfy import samplers
|
||||
|
||||
@@ -55,8 +58,8 @@ class NoisyLatentLikeNode:
|
||||
latent["samples"],
|
||||
None,
|
||||
None,
|
||||
seed=None,
|
||||
use_cpu=True,
|
||||
seed=seed,
|
||||
cpu=True,
|
||||
)
|
||||
randst = torch.random.get_rng_state()
|
||||
try:
|
||||
@@ -113,7 +116,7 @@ class SonarCustomNoiseNode:
|
||||
nis = (
|
||||
sonar_custom_noise_opt.clone()
|
||||
if sonar_custom_noise_opt
|
||||
else noise.CustomNoise()
|
||||
else noise.CustomNoiseChain()
|
||||
)
|
||||
if factor != 0:
|
||||
nis.add(noise.CustomNoiseItem(factor, noise_type))
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
+143
-52
@@ -1,7 +1,9 @@
|
||||
# Noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers
|
||||
from __future__ import annotations
|
||||
|
||||
import functools as fun
|
||||
import math
|
||||
import operator as op
|
||||
from enum import Enum, auto
|
||||
from typing import Callable
|
||||
|
||||
@@ -12,13 +14,18 @@ from torch import FloatTensor, Generator, Tensor
|
||||
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
class NoiseType(Enum):
|
||||
GAUSSIAN = auto()
|
||||
UNIFORM = auto()
|
||||
BROWNIAN = auto()
|
||||
PERLIN = auto()
|
||||
STUDENTT = auto()
|
||||
STUDENTT_TEST = auto()
|
||||
HIGHRES_PYRAMID = auto()
|
||||
PINK = auto()
|
||||
LAPLACIAN = auto()
|
||||
@@ -41,12 +48,12 @@ class CustomNoiseItem:
|
||||
self.noise_type = noise_type
|
||||
|
||||
|
||||
class CustomNoise:
|
||||
class CustomNoiseChain:
|
||||
def __init__(self, items=None):
|
||||
self.items = items if items is not None else []
|
||||
|
||||
def clone(self):
|
||||
return CustomNoise(
|
||||
return CustomNoiseChain(
|
||||
[CustomNoiseItem(i.factor, i.noise_type) for i in self.items],
|
||||
)
|
||||
|
||||
@@ -56,27 +63,53 @@ class CustomNoise:
|
||||
def rescaled(self, scale=1.0):
|
||||
total = sum(i.factor for i in self.items)
|
||||
divisor = total / scale
|
||||
return CustomNoise(
|
||||
divisor = divisor if divisor != 0 else 1.0
|
||||
return CustomNoiseChain(
|
||||
[CustomNoiseItem(i.factor / divisor, i.noise_type) 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, x: Tensor) -> Callable:
|
||||
items = tuple(
|
||||
(get_noise_sampler(i.noise_type, x, None, None), i.factor)
|
||||
def make_noise_sampler(
|
||||
self,
|
||||
x: Tensor,
|
||||
sigma_min=None,
|
||||
sigma_max=None,
|
||||
seed=None,
|
||||
cpu=True,
|
||||
) -> Callable:
|
||||
noise_samplers = tuple(
|
||||
get_noise_sampler(
|
||||
i.noise_type,
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
factor=i.factor,
|
||||
)
|
||||
for i in self.items
|
||||
)
|
||||
if not items or not all(i[0] for i in items):
|
||||
if not noise_samplers or not all(noise_samplers):
|
||||
raise ValueError("Failed to get noise sampler")
|
||||
scale = sum(i.factor for i in self.items)
|
||||
|
||||
def noise_sampler(s, sn):
|
||||
nonlocal items
|
||||
result = items[0][0](s, sn) * items[0][1]
|
||||
for ns, factor in items[1:]:
|
||||
result += ns(s, sn) * factor
|
||||
result /= result.std()
|
||||
scale = sum(i[1] for i in items)
|
||||
return result * scale
|
||||
def noise_sampler(sigma, sigma_next):
|
||||
result = fun.reduce(
|
||||
op.add,
|
||||
(ns(sigma, sigma_next) for ns in noise_samplers),
|
||||
)
|
||||
return scale_noise(result, scale)
|
||||
|
||||
return noise_sampler
|
||||
|
||||
@@ -369,8 +402,9 @@ def power_noise_like(tensor, alpha=2, k=1): # This doesn't work properly right
|
||||
"""
|
||||
tensor = torch.randn_like(tensor)
|
||||
fft = torch.fft.fft2(tensor)
|
||||
freq = torch.arange(1, len(fft) + 1, dtype=torch.float)
|
||||
freq = freq.reshape(freq.shape + (1,) * (len(tensor.shape) - 1))
|
||||
freq = torch.arange(1, len(fft) + 1, dtype=torch.float).reshape(
|
||||
(len(fft),) + (1,) * (tensor.dim() - 1),
|
||||
)
|
||||
spectral_density = k / freq**alpha
|
||||
noise = torch.rand(tensor.shape) * spectral_density
|
||||
mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
|
||||
@@ -378,28 +412,83 @@ def power_noise_like(tensor, alpha=2, k=1): # This doesn't work properly right
|
||||
return noise.to(tensor.device).sub_(mean).div_(std)
|
||||
|
||||
|
||||
class NoiseSampler:
|
||||
def __init__(
|
||||
self,
|
||||
x: Tensor,
|
||||
sigma_min: float | None = None,
|
||||
sigma_max: float | None = None,
|
||||
seed: int | None = None,
|
||||
cpu: bool = False,
|
||||
transform: Callable = lambda t: t,
|
||||
make_noise_sampler: Callable | None = None,
|
||||
normalize_noise=False,
|
||||
factor: float = 1.0,
|
||||
):
|
||||
try:
|
||||
self.noise_sampler = make_noise_sampler(
|
||||
x,
|
||||
transform(torch.as_tensor(sigma_min))
|
||||
if sigma_min is not None
|
||||
else None,
|
||||
transform(torch.as_tensor(sigma_max))
|
||||
if sigma_max is not None
|
||||
else None,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
)
|
||||
except TypeError:
|
||||
self.noise_sampler = make_noise_sampler(x)
|
||||
self.factor = factor
|
||||
self.normalize_noise = normalize_noise
|
||||
self.transform = transform
|
||||
self.device = x.device
|
||||
self.dtype = x.dtype
|
||||
|
||||
@classmethod
|
||||
def simple(cls, f):
|
||||
return lambda *args, **kwargs: cls(
|
||||
*args,
|
||||
**kwargs,
|
||||
make_noise_sampler=lambda x, *_args, **_kwargs: lambda _s, _sn: f(x),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def wrap(cls, f):
|
||||
return lambda *args, **kwargs: cls(*args, **kwargs, make_noise_sampler=f)
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
args = (
|
||||
self.transform(torch.as_tensor(s)) if s is not None else s for s in args
|
||||
)
|
||||
noise = self.noise_sampler(*args, **kwargs)
|
||||
noise = (
|
||||
scale_noise(noise, self.factor)
|
||||
if self.normalize_noise
|
||||
else noise.mul_(self.factor)
|
||||
)
|
||||
if hasattr(noise, "to"):
|
||||
return noise.to(dtype=self.dtype, device=self.device)
|
||||
return noise
|
||||
|
||||
|
||||
NOISE_SAMPLERS: dict[NoiseType, Callable] = {
|
||||
# No brownian as it is a special case that requires extra stuff like seed.
|
||||
NoiseType.GAUSSIAN: sampling.default_noise_sampler,
|
||||
NoiseType.UNIFORM: lambda x: lambda _s, _sn: uniform_noise_like(x),
|
||||
NoiseType.PERLIN: lambda x: lambda _s, _sn: rand_perlin_like(x),
|
||||
NoiseType.STUDENTT: studentt_noise_sampler,
|
||||
NoiseType.STUDENTT_TEST: lambda x: lambda _s, _sn: studentt_noise_like(x).to(
|
||||
x.device,
|
||||
NoiseType.BROWNIAN: NoiseSampler.wrap(sampling.BrownianTreeNoiseSampler),
|
||||
NoiseType.GAUSSIAN: NoiseSampler.simple(torch.randn_like),
|
||||
NoiseType.UNIFORM: NoiseSampler.simple(uniform_noise_like),
|
||||
NoiseType.PERLIN: NoiseSampler.simple(rand_perlin_like),
|
||||
NoiseType.STUDENTT: NoiseSampler.simple(studentt_noise_like),
|
||||
NoiseType.PINK: NoiseSampler.simple(pink_noise_like),
|
||||
NoiseType.HIGHRES_PYRAMID: NoiseSampler.simple(highres_pyramid_noise_like),
|
||||
NoiseType.RAINBOW_MILD: NoiseSampler.simple(
|
||||
lambda x: (green_noise_like(x) * 0.55 + rand_perlin_like(x) * 0.7) * 1.15,
|
||||
),
|
||||
NoiseType.PINK: lambda x: lambda _s, _sn: pink_noise_like(x),
|
||||
NoiseType.HIGHRES_PYRAMID: lambda x: lambda _s, _sn: highres_pyramid_noise_like(x),
|
||||
NoiseType.RAINBOW_MILD: lambda x: lambda _s, _sn: (
|
||||
green_noise_like(x) * 0.55 + rand_perlin_like(x) * 0.7
|
||||
)
|
||||
* 1.15,
|
||||
NoiseType.RAINBOW_INTENSE: lambda x: lambda _s, _sn: (
|
||||
green_noise_like(x) * 0.75 + rand_perlin_like(x) * 0.5
|
||||
)
|
||||
* 1.15,
|
||||
NoiseType.LAPLACIAN: lambda x: lambda _s, _sn: laplacian_noise_like(x),
|
||||
NoiseType.POWER: lambda x: lambda _s, _sn: power_noise_like(x),
|
||||
NoiseType.GREEN_TEST: lambda x: lambda _s, _sn: green_noise_like(x),
|
||||
NoiseType.RAINBOW_INTENSE: NoiseSampler.simple(
|
||||
lambda x: (green_noise_like(x) * 0.75 + rand_perlin_like(x) * 0.5) * 1.15,
|
||||
),
|
||||
NoiseType.LAPLACIAN: NoiseSampler.simple(laplacian_noise_like),
|
||||
NoiseType.POWER: NoiseSampler.simple(power_noise_like),
|
||||
NoiseType.GREEN_TEST: NoiseSampler.simple(green_noise_like),
|
||||
# NoiseType.RAINBOW_MILD2: lambda x: lambda _s, _sn: (
|
||||
# green_noise_like(x) * 0.55 + uniform_noise_like(x) * 0.7
|
||||
# )
|
||||
@@ -421,23 +510,25 @@ def get_noise_sampler(
|
||||
sigma_min: float | None,
|
||||
sigma_max: float | None,
|
||||
seed: int | None = None,
|
||||
use_cpu: bool = True,
|
||||
cpu: bool = True,
|
||||
factor: float = 1.0,
|
||||
normalize_noise=True,
|
||||
) -> Callable:
|
||||
if noise_type is None:
|
||||
noise_type = NoiseType.GAUSSIAN
|
||||
elif isinstance(noise_type, str):
|
||||
noise_type = NoiseType[noise_type.upper()]
|
||||
if noise_type == NoiseType.BROWNIAN:
|
||||
if sigma_min is None or sigma_max is None:
|
||||
raise ValueError("Must pass sigma min/max when using brownian noise")
|
||||
return sampling.BrownianTreeNoiseSampler(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=use_cpu,
|
||||
)
|
||||
ns = NOISE_SAMPLERS.get(noise_type)
|
||||
if ns is None:
|
||||
if noise_type == NoiseType.BROWNIAN and (sigma_min is None or sigma_max is None):
|
||||
raise ValueError("Must pass sigma min/max when using brownian noise")
|
||||
mkns = NOISE_SAMPLERS.get(noise_type)
|
||||
if mkns is None:
|
||||
raise ValueError("Unknown noise sampler")
|
||||
return ns(x)
|
||||
return mkns(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
factor=factor,
|
||||
normalize_noise=normalize_noise,
|
||||
)
|
||||
|
||||
+65
-40
@@ -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
|
||||
@@ -67,7 +99,7 @@ class SonarBase:
|
||||
None,
|
||||
None,
|
||||
seed=self.extra_args.get("seed"),
|
||||
use_cpu=True,
|
||||
cpu=True,
|
||||
)
|
||||
self.history_d = ns(None, None)
|
||||
else:
|
||||
@@ -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.shape)
|
||||
|
||||
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,12 @@ 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"),
|
||||
)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
x, sigma, sigma_hat, denoised = sonar.step(
|
||||
@@ -285,14 +327,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 +382,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,21 +394,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()
|
||||
if sonar_config.custom_noise:
|
||||
noise_sampler = sonar_config.custom_noise.make_noise_sampler(x)
|
||||
else:
|
||||
noise_sampler = noise.get_noise_sampler(
|
||||
sonar_config.noise_type,
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=extra_args.get("seed"),
|
||||
use_cpu=True,
|
||||
)
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sonar = cls(
|
||||
noise_sampler,
|
||||
eta,
|
||||
s_noise,
|
||||
model,
|
||||
@@ -377,6 +404,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(
|
||||
@@ -399,14 +432,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
|
||||
|
||||
@@ -537,21 +568,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()
|
||||
if sonar_config.custom_noise:
|
||||
noise_sampler = sonar_config.custom_noise.make_noise_sampler(x)
|
||||
else:
|
||||
noise_sampler = noise.get_noise_sampler(
|
||||
sonar_config.noise_type,
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=extra_args.get("seed"),
|
||||
use_cpu=True,
|
||||
)
|
||||
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sonar = cls(
|
||||
noise_sampler,
|
||||
eta,
|
||||
s_noise,
|
||||
model,
|
||||
@@ -560,6 +579,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(
|
||||
|
||||
Reference in New Issue
Block a user