Refactor noise: stage 3
This commit is contained in:
@@ -164,7 +164,7 @@ Normal (non-sonar) Eular A. Not really a comparison with noise (think it would u
|
||||
|
||||

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

|
||||
|
||||
|
||||
+8
-2
@@ -16,7 +16,9 @@ from torch import FloatTensor, Generator, Tensor
|
||||
|
||||
# This likely isn't correct.
|
||||
def scale_noise(noise, factor=1.0):
|
||||
return noise * factor
|
||||
mean, std = noise.mean(), noise.std()
|
||||
# print(f"factor={factor:.3}, mean={mean:.3}, std={std:.3}")
|
||||
noise = noise - mean
|
||||
if std >= 0.98:
|
||||
noise /= std
|
||||
@@ -115,7 +117,9 @@ class CustomNoiseChain:
|
||||
op.add,
|
||||
(ns(sigma, sigma_next) for ns in noise_samplers),
|
||||
)
|
||||
return scale_noise(result, scale)
|
||||
result = scale_noise(result, scale)
|
||||
# print("SCALING", scale, "sum", result.sum().item())
|
||||
return result
|
||||
|
||||
return noise_sampler
|
||||
|
||||
@@ -465,7 +469,9 @@ class NoiseSampler:
|
||||
args = (
|
||||
self.transform(torch.as_tensor(s)) if s is not None else s for s in args
|
||||
)
|
||||
result = scale_noise(self.noise_sampler(*args, **kwargs), self.factor)
|
||||
# print(">", self.factor)
|
||||
result = self.noise_sampler(*args, **kwargs) * 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
|
||||
|
||||
+65
-51
@@ -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
|
||||
@@ -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)
|
||||
|
||||
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,13 @@ 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"),
|
||||
)
|
||||
# print("noise sampler:", noise_sampler, sonar.noise_sampler)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
x, sigma, sigma_hat, denoised = sonar.step(
|
||||
@@ -285,14 +328,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 +383,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,27 +395,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()
|
||||
seed = extra_args.get("seed")
|
||||
if sonar_config.custom_noise:
|
||||
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=seed,
|
||||
cpu=True,
|
||||
)
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sonar = cls(
|
||||
noise_sampler,
|
||||
eta,
|
||||
s_noise,
|
||||
model,
|
||||
@@ -383,6 +405,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(
|
||||
@@ -405,14 +433,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
|
||||
|
||||
@@ -543,27 +569,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()
|
||||
seed = extra_args.get("seed")
|
||||
if sonar_config.custom_noise:
|
||||
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=seed,
|
||||
cpu=True,
|
||||
)
|
||||
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sonar = cls(
|
||||
noise_sampler,
|
||||
eta,
|
||||
s_noise,
|
||||
model,
|
||||
@@ -572,6 +580,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