Momentum sampler refactor/improvements (I hope)
This commit is contained in:
+41
-48
@@ -1703,7 +1703,7 @@ class SamplerNodeSonarBase:
|
||||
"max": 2.5,
|
||||
"step": 0.01,
|
||||
"round": False,
|
||||
"tooltip": "Strength of the output from normal sampling. When set to 1.0 effectively disables momentum.",
|
||||
"tooltip": "How much of the normal result to keep during sampling. 0.95 means 95% normal, 5% from history. When set to 1.0 effectively disables momentum.",
|
||||
},
|
||||
),
|
||||
"momentum_hist": (
|
||||
@@ -1714,7 +1714,7 @@ class SamplerNodeSonarBase:
|
||||
"max": 1.5,
|
||||
"step": 0.01,
|
||||
"round": False,
|
||||
"tooltip": "Strength of momentum history",
|
||||
"tooltip": "How much of the existing history to leave at each update. 0.75 means keep 75%, mix in 25% of the new result.",
|
||||
},
|
||||
),
|
||||
"momentum_init": (
|
||||
@@ -1756,6 +1756,34 @@ class SamplerNodeSonarBase:
|
||||
|
||||
|
||||
class SamplerNodeSonarEuler(SamplerNodeSonarBase):
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
@classmethod
|
||||
def get_sampler(
|
||||
cls,
|
||||
*,
|
||||
momentum,
|
||||
momentum_hist,
|
||||
momentum_init,
|
||||
direction,
|
||||
rand_init_noise_type,
|
||||
guidance_cfg_opt=None,
|
||||
):
|
||||
cfg = SonarConfig(
|
||||
momentum=momentum,
|
||||
init=HistoryType[momentum_init.upper()],
|
||||
momentum_hist=momentum_hist,
|
||||
direction=direction,
|
||||
rand_init_noise_type=NoiseType[rand_init_noise_type.upper()],
|
||||
guidance=guidance_cfg_opt,
|
||||
)
|
||||
return (samplers.KSAMPLER(SonarEuler.sampler, {"sonar_config": cfg}),)
|
||||
|
||||
|
||||
class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
result = super().INPUT_TYPES()
|
||||
@@ -1772,52 +1800,6 @@ class SamplerNodeSonarEuler(SamplerNodeSonarBase):
|
||||
"tooltip": "Multiplier for noise added during ancestral or SDE sampling.",
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
return result
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
@classmethod
|
||||
def get_sampler(
|
||||
cls,
|
||||
*,
|
||||
momentum,
|
||||
momentum_hist,
|
||||
momentum_init,
|
||||
direction,
|
||||
rand_init_noise_type,
|
||||
s_noise,
|
||||
guidance_cfg_opt=None,
|
||||
):
|
||||
cfg = SonarConfig(
|
||||
momentum=momentum,
|
||||
init=HistoryType[momentum_init.upper()],
|
||||
momentum_hist=momentum_hist,
|
||||
direction=direction,
|
||||
rand_init_noise_type=NoiseType[rand_init_noise_type.upper()],
|
||||
guidance=guidance_cfg_opt,
|
||||
)
|
||||
return (
|
||||
samplers.KSAMPLER(
|
||||
SonarEuler.sampler,
|
||||
{
|
||||
"s_noise": s_noise,
|
||||
"sonar_config": cfg,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
result = super().INPUT_TYPES()
|
||||
result["required"].update(
|
||||
{
|
||||
"eta": (
|
||||
"FLOAT",
|
||||
{
|
||||
@@ -1892,6 +1874,17 @@ class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEuler):
|
||||
result = super().INPUT_TYPES()
|
||||
result["required"].update(
|
||||
{
|
||||
"s_noise": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.01,
|
||||
"round": False,
|
||||
"tooltip": "Multiplier for noise added during ancestral or SDE sampling.",
|
||||
},
|
||||
),
|
||||
"eta": (
|
||||
"FLOAT",
|
||||
{
|
||||
|
||||
+65
-67
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from enum import Enum, auto
|
||||
from functools import lru_cache
|
||||
from sys import stderr
|
||||
from typing import Any, Callable, NamedTuple
|
||||
|
||||
@@ -14,6 +15,7 @@ from torch import Tensor
|
||||
from tqdm.auto import trange
|
||||
|
||||
from . import noise
|
||||
from .noise_utils import BLENDING_MODES
|
||||
|
||||
|
||||
class HistoryType(Enum):
|
||||
@@ -44,15 +46,17 @@ class SonarConfig(NamedTuple):
|
||||
custom_noise: noise.CustomNoise | None = None
|
||||
rand_init_noise_type: noise.NoiseType | None = None
|
||||
guidance: GuidanceConfig | None = None
|
||||
blend_mode: str = "lerp"
|
||||
|
||||
|
||||
class SonarBase:
|
||||
DEFAULT_NOISE_TYPE = noise.NoiseType.GAUSSIAN
|
||||
|
||||
def __init__(self, cfg: SonarConfig) -> None:
|
||||
def __init__(self, cfg: SonarConfig, *, blend_mode="lerp") -> None:
|
||||
self.history_d = None
|
||||
self.cfg = cfg
|
||||
self.noise_sampler = None
|
||||
self.blend_function = BLENDING_MODES[blend_mode]
|
||||
|
||||
def set_noise_sampler(
|
||||
self,
|
||||
@@ -95,7 +99,7 @@ class SonarBase:
|
||||
return
|
||||
# memorize delta momentum
|
||||
if self.cfg.init == HistoryType.ZERO:
|
||||
self.history_d = 0
|
||||
self.history_d = None
|
||||
elif self.cfg.init == HistoryType.SAMPLE:
|
||||
self.history_d = x
|
||||
elif self.cfg.init == HistoryType.RAND:
|
||||
@@ -112,21 +116,37 @@ class SonarBase:
|
||||
else:
|
||||
raise ValueError("Sonar sampler: bad history type")
|
||||
|
||||
def update_hist(self, momentum_d):
|
||||
q = 1.0 - self.cfg.momentum_hist
|
||||
hd = self.history_d
|
||||
if isinstance(hd, int) and hd == 0:
|
||||
self.history_d = momentum_d
|
||||
else:
|
||||
self.history_d = (1.0 - q) * hd + q * momentum_d
|
||||
@property
|
||||
@lru_cache(maxsize=1) # noqa: B019
|
||||
def history_ratios(self):
|
||||
direction = self.cfg.direction
|
||||
momentum_hist = self.cfg.momentum_hist
|
||||
return (
|
||||
momentum_hist,
|
||||
1.0 + abs(direction) * (1 - momentum_hist)
|
||||
if direction < 0
|
||||
else 2.0 - direction,
|
||||
direction,
|
||||
)
|
||||
|
||||
def update_hist(self, momentum_d: torch.Tensor) -> None:
|
||||
hd, cfg = self.history_d, self.cfg
|
||||
if cfg.momentum_hist == 1:
|
||||
return
|
||||
hd_ratio, hd_scale, md_scale = self.history_ratios
|
||||
self.history_d = (
|
||||
momentum_d
|
||||
if hd is None
|
||||
else self.blend_function(momentum_d * md_scale, hd * hd_scale, hd_ratio)
|
||||
)
|
||||
|
||||
def momentum_step(self, x: Tensor, d: Tensor, dt: Tensor):
|
||||
if self.cfg.momentum == 1.0:
|
||||
momentum = self.cfg.momentum
|
||||
if momentum == 1.0:
|
||||
return x + d * dt
|
||||
hd = self.history_d
|
||||
# correct current `d` with momentum
|
||||
p = (1.0 - self.cfg.momentum) * self.cfg.direction
|
||||
momentum_d = (1.0 - p) * d + p * hd
|
||||
momentum_d = d if hd is None else self.blend_function(hd, d, momentum)
|
||||
|
||||
# Euler method with momentum
|
||||
x = x + momentum_d * dt # noqa: PLR6104
|
||||
@@ -152,8 +172,8 @@ class SonarGuidanceMixin:
|
||||
def prepare_ref_latent(latent: Tensor | None) -> Tensor:
|
||||
if latent is None:
|
||||
return None
|
||||
avg_s = latent.mean(dim=[2, 3], keepdim=True)
|
||||
std_s = latent.std(dim=[2, 3], keepdim=True)
|
||||
avg_s = latent.mean(dim=(-2, -1), keepdim=True)
|
||||
std_s = latent.std(dim=(-2, -1), keepdim=True)
|
||||
return ((latent - avg_s) / std_s).to(latent.dtype)
|
||||
|
||||
def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor):
|
||||
@@ -188,8 +208,8 @@ class SonarGuidanceMixin:
|
||||
ref_latent: Tensor,
|
||||
factor: float = 0.2,
|
||||
) -> Tensor:
|
||||
avg_t = denoised.mean(dim=[1, 2, 3], keepdim=True)
|
||||
std_t = denoised.std(dim=[1, 2, 3], keepdim=True)
|
||||
avg_t = denoised.mean(dim=(-3, -2, -1), keepdim=True)
|
||||
std_t = denoised.std(dim=(-3, -2, -1), keepdim=True)
|
||||
ref_img_shift = ref_latent * std_t + avg_t
|
||||
|
||||
d = sampling.to_d(x, sigma, ref_img_shift)
|
||||
@@ -198,8 +218,8 @@ class SonarGuidanceMixin:
|
||||
|
||||
@staticmethod
|
||||
def guidance_linear(x: Tensor, ref_latent: Tensor, factor: float = 0.2) -> Tensor:
|
||||
avg_t = x.mean(dim=[1, 2, 3], keepdim=True)
|
||||
std_t = x.std(dim=[1, 2, 3], keepdim=True)
|
||||
avg_t = x.mean(dim=(-3, -2, -1), keepdim=True)
|
||||
std_t = x.std(dim=(-3, -2, -1), keepdim=True)
|
||||
ref_img_shift = ref_latent * std_t + avg_t
|
||||
return (1.0 - factor) * x + factor * ref_img_shift
|
||||
|
||||
@@ -230,58 +250,28 @@ class SonarSampler(SonarWithGuidance):
|
||||
class SonarEuler(SonarSampler):
|
||||
def __init__(
|
||||
self,
|
||||
s_churn: float = 0.0,
|
||||
s_tmin: float = 0.0,
|
||||
s_tmax: float = float("inf"),
|
||||
s_noise: float = 1.0,
|
||||
*args: list[Any],
|
||||
**kwargs: dict[str, Any],
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.s_churn = s_churn
|
||||
self.s_tmin = s_tmin
|
||||
self.s_tmax = s_tmax
|
||||
self.s_noise = s_noise
|
||||
|
||||
def step(
|
||||
self,
|
||||
step_index: int,
|
||||
sample: torch.FloatTensor,
|
||||
):
|
||||
def step(self, step_index: int, sample: torch.FloatTensor):
|
||||
self.init_hist_d(sample)
|
||||
|
||||
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)
|
||||
if self.s_tmin <= sigma <= self.s_tmax
|
||||
else 0.0
|
||||
)
|
||||
|
||||
sigma_hat = sigma * (gamma + 1)
|
||||
|
||||
if gamma > 0:
|
||||
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 # noqa: PLR6104
|
||||
|
||||
denoised = self.model(sample, sigma_hat * self.s_in, **self.extra_args)
|
||||
denoised = self.model(sample, sigma * self.s_in, **self.extra_args)
|
||||
derivative = sampling.to_d(sample, sigma, denoised)
|
||||
dt = self.sigmas[step_index + 1] - sigma_hat
|
||||
dt = sigma_to - sigma
|
||||
|
||||
result_sample = self.momentum_step(sample, derivative, dt)
|
||||
|
||||
if self.sigmas[step_index + 1] > 0:
|
||||
if sigma_to > 0:
|
||||
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
||||
|
||||
return (
|
||||
result_sample,
|
||||
sigma,
|
||||
sigma_hat,
|
||||
sigma,
|
||||
denoised,
|
||||
)
|
||||
|
||||
@@ -296,25 +286,19 @@ class SonarEuler(SonarSampler):
|
||||
callback=None,
|
||||
disable=None,
|
||||
noise_sampler: Callable | None = None,
|
||||
sonar_blend_mode="lerp",
|
||||
sonar_config=None,
|
||||
s_churn=0.0,
|
||||
s_tmin=0.0,
|
||||
s_tmax=float("inf"),
|
||||
s_noise=1.0,
|
||||
):
|
||||
if sonar_config is None:
|
||||
sonar_config = SonarConfig()
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
s_in = x.new_ones((x.shape[0],))
|
||||
sonar = cls(
|
||||
s_churn,
|
||||
s_tmin,
|
||||
s_tmax,
|
||||
s_noise,
|
||||
model,
|
||||
sigmas,
|
||||
s_in,
|
||||
{} if extra_args is None else extra_args,
|
||||
sonar_config,
|
||||
blend_mode=sonar_blend_mode,
|
||||
)
|
||||
sonar.set_noise_sampler(
|
||||
x,
|
||||
@@ -324,7 +308,7 @@ class SonarEuler(SonarSampler):
|
||||
)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
x, _sigma, sigma_hat, denoised = sonar.step(
|
||||
x, sigma, sigma_hat, denoised = sonar.step(
|
||||
i,
|
||||
x,
|
||||
)
|
||||
@@ -333,7 +317,7 @@ class SonarEuler(SonarSampler):
|
||||
{
|
||||
"x": x,
|
||||
"i": i,
|
||||
"sigma": sigmas[i],
|
||||
"sigma": sigma,
|
||||
"sigma_hat": sigma_hat,
|
||||
"denoised": denoised,
|
||||
},
|
||||
@@ -396,6 +380,7 @@ class SonarEulerAncestral(SonarSampler):
|
||||
extra_args=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
sonar_blend_mode="lerp",
|
||||
sonar_config=None,
|
||||
eta=1.0,
|
||||
s_noise=1.0,
|
||||
@@ -412,6 +397,7 @@ class SonarEulerAncestral(SonarSampler):
|
||||
s_in,
|
||||
{} if extra_args is None else extra_args,
|
||||
sonar_config,
|
||||
blend_mode=sonar_blend_mode,
|
||||
)
|
||||
sonar.set_noise_sampler(
|
||||
x,
|
||||
@@ -482,7 +468,12 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
return sigma.log().neg()
|
||||
|
||||
hd = self.history_d
|
||||
p = (1.0 - self.cfg.momentum) * self.cfg.direction
|
||||
# Halve the momentum proportion if there's history since we will use it twice.
|
||||
adjusted_momentum = (
|
||||
self.cfg.momentum + (1 - self.cfg.momentum) / 2
|
||||
if hd is not None
|
||||
else self.cfg.momentum
|
||||
)
|
||||
|
||||
r = 1 / 2
|
||||
# DPM-Solver++
|
||||
@@ -495,7 +486,9 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
sd, su = sampling.get_ancestral_step(sigma_fn(t), sigma_fn(s), self.eta)
|
||||
s_ = t_fn(sd)
|
||||
diff_2 = (t - s_).expm1() * denoised
|
||||
momentum_d = (1.0 - p) * diff_2 + p * hd
|
||||
momentum_d = (
|
||||
diff_2 if hd is None else self.blend_function(hd, diff_2, adjusted_momentum)
|
||||
)
|
||||
self.update_hist(momentum_d)
|
||||
hd = self.history_d
|
||||
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - momentum_d
|
||||
@@ -511,7 +504,10 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
t_next_ = t_fn(sd)
|
||||
denoised_d = (1 - fac) * denoised + fac * denoised_2
|
||||
diff_1 = (t - t_next_).expm1() * denoised_d
|
||||
momentum_d = (1.0 - p) * diff_1 + p * hd
|
||||
hd = self.history_d
|
||||
momentum_d = (
|
||||
diff_1 if hd is None else self.blend_function(hd, diff_1, adjusted_momentum)
|
||||
)
|
||||
self.update_hist(momentum_d)
|
||||
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - momentum_d
|
||||
x = self.guidance_step(step_index, x, denoised_d)
|
||||
@@ -564,6 +560,7 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
extra_args=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
sonar_blend_mode="lerp",
|
||||
sonar_config=None,
|
||||
eta=1.0,
|
||||
s_noise=1.0,
|
||||
@@ -580,6 +577,7 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
s_in,
|
||||
{} if extra_args is None else extra_args,
|
||||
sonar_config,
|
||||
blend_mode=sonar_blend_mode,
|
||||
)
|
||||
sonar.set_noise_sampler(
|
||||
x,
|
||||
|
||||
Reference in New Issue
Block a user