Refactor noise: stage 2

This commit is contained in:
blepping
2024-02-24 09:59:45 -07:00
parent 8f2632a8fe
commit 227be2c85a
5 changed files with 190 additions and 15 deletions
+1
View File
@@ -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,
+5
View File
@@ -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.
+151 -1
View File
@@ -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,
)
+14 -7
View File
@@ -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
+19 -7
View File
@@ -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]])