Refactor noise: stage 2
This commit is contained in:
@@ -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,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
@@ -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
@@ -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
@@ -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]])
|
||||
|
||||
Reference in New Issue
Block a user