Add DPM++ 3M SDE DynETA

Partially implement PR #4 & #5 from blepping (oops hold on)
This commit is contained in:
Clybius
2024-03-04 12:32:26 -06:00
parent f3646e3790
commit 1f161b20f9
3 changed files with 154 additions and 71 deletions
+3
View File
@@ -14,5 +14,8 @@ NODE_CLASS_MAPPINGS = {
"SamplerTTM": nodes.SamplerTTM,
"SamplerLCMCustom": nodes.SamplerLCMCustom,
"SamplerEulerAncestralDancing_Experimental": nodes.SamplerEULER_ANCESTRAL_DANCING,
"SamplerDPMPP_3M_SDE_DynETA": nodes.SamplerDPMPP_3M_SDE_DYN_ETA,
### Schedulers
"SimpleExponentialScheduler": nodes.SimpleExponentialScheduler,
}
__all__ = ['NODE_CLASS_MAPPINGS']
+102 -67
View File
@@ -8,7 +8,6 @@ from tqdm.auto import trange, tqdm
import comfy.sample
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, PIDStepSizeController, get_ancestral_step, to_d, default_noise_sampler
import random
# The following function adds the samplers during initialization, in __init__.py
def add_samplers():
@@ -32,16 +31,19 @@ def add_samplers():
# The following function adds the samplers during initialization, in __init__.py
def add_schedulers():
from comfy.samplers import KSampler, k_diffusion_sampling
added = 0
for scheduler in extra_schedulers: #getattr(self, "sample_{}".format(extra_samplers))
if scheduler not in KSampler.SCHEDULERS:
try:
idx = KSampler.SCHEDULERS.index("ddim_uniform") # Last item in the samplers list
KSampler.SCHEDULERS.insert(idx+1, scheduler) # Add our custom samplers
setattr(k_diffusion_sampling, "get_sigmas_{}".format(scheduler), extra_schedulers[scheduler])
import importlib
importlib.reload(k_diffusion_sampling)
added += 1
except ValueError as err:
pass
if added > 0:
import importlib
importlib.reload(k_diffusion_sampling)
# Noise samplers
from torch import Generator, Tensor, lerp
@@ -264,7 +266,7 @@ def highres_pyramid_noise_like(x, discount=0.7):
u = torch.nn.Upsample(size=(orig_h, orig_w), mode='bilinear')
noise = (torch.rand_like(x) - 0.5) * 2 * 1.73 # Start with scaled uniform noise
for i in range(4):
r = random.random()*2+2 # Rather than always going 2x,
r = torch.rand(1).item() * 2 + 2 # Rather than always going 2x,
h, w = min(orig_h*15, int(h*(r**i))), min(orig_w*15, int(w*(r**i)))
noise += u(torch.randn(b, c, h, w).to(x)) * discount**i
if h>=orig_h*15 or w>=orig_w*15: break # Lowest resolution is 1x1
@@ -648,7 +650,7 @@ def sample_clyb_4m_sde(model, x, sigmas, extra_args=None, callback=None, disable
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
return sample_clyb_4m_sde_momentumized(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, momentum=momentum)
"""
# This code works, but I'm currently experimenting with different methods
@torch.no_grad()
def sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, leap=2, eta_dance=1.0):
@@ -695,67 +697,6 @@ def sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
return x
"""
def rej(a, b):
"""
Implements the rejection function for alternative diffusion.
Args:
a: Tensor of shape (B, H, W, C), where B is batch size, H and W are spatial dimensions, and C is number of channels.
b: Tensor of the same shape as a.
Returns:
Tensor of the same shape as a and b, containing the rejection output.
"""
return (b * torch.tensordot(a, b, dims=len(a.shape)) / torch.tensordot(b, b, dims=len(a.shape))) - a
@torch.no_grad()
def sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, leap=2, eta_dance=1.0):
#Ancestral sampling with Euler method steps, dancing steps.
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
unsample_noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
s_in = x.new_ones([x.shape[0]])
for i in trange(len(sigmas) - 1, disable=disable):
if i < len(sigmas) - leap:
is_danceable = sigmas[i + leap] > 0
else:
is_danceable = False
orig_x = x
denoised = model(x, sigmas[i] * s_in, **extra_args)
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + leap] if is_danceable else sigmas[i + 1], eta=eta)
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
d = to_d(x, sigmas[i], denoised)
# Euler method
dt = sigma_down - sigmas[i]
x = x + d * dt
if sigmas[i + 1] > 0:
if is_danceable:
#x = x + noise_sampler(sigmas[i], sigmas[i + leap]) * s_noise * sigma_up
#x = x + noise_sampler(sigmas[i + 2], sigmas[i + 1]) * s_noise * sigma_up
denoised2 = model(x, sigmas[i + leap] * s_in, **extra_args)
sigma_down2, sigma_up2 = get_ancestral_step(sigmas[i + leap], sigmas[i + 1], eta=eta_dance)
d_2 = to_d(x, sigmas[i + leap], denoised2)
dt_2 = sigma_down2 - sigmas[i + leap]
x_2 = x + d_2 * dt_2
#x_2 = x_2 + noise_sampler(sigmas[i + leap], sigmas[i]) * s_noise * sigma_up2
sigma_down3, sigma_up3 = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
#x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up3
denoised3 = model(x_2, sigmas[i] * s_in, **extra_args)
d_3 = to_d(orig_x, sigmas[i], denoised3 + rej(denoised3 - denoised, denoised2 - denoised))
#d_3 = to_d(x_2, sigmas[i], denoised3)
dt_3 = sigma_down3 - sigmas[i]
x = orig_x + d_3 * dt_3 # Very denoised, slightly denoised
#print(dt_3, dt_2)
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up3
#x = x + d * dt
else:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
return x
def sample_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler="gaussian", leap=2, eta_dance=1.0):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
@@ -777,6 +718,86 @@ def sample_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=N
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
return sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, leap=leap, eta_dance=eta_dance)
@torch.no_grad()
def sampler_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback=None, disable=None, eta_max=1.0, eta_min=0.0, s_noise=1., noise_sampler=None):
"""DPM-Solver++(3M) SDE with dynamic eta."""
def eta_schedule_cosine_annealing(i, n, eta_max=eta_max, eta_min=eta_min):
"""Cosine annealing schedule for eta."""
progress = i / (n - 1)
eta = eta_min + 0.5 * (eta_max - eta_min) * (1 + math.cos(math.pi * progress))
return eta
seed = extra_args.get("seed", None)
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) if noise_sampler is None else noise_sampler
extra_args = {} if extra_args is None else extra_args
s_in = x.new_ones([x.shape[0]])
denoised_1, denoised_2 = None, None
h, h_1, h_2 = None, None, None
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args)
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
# DPM-Solver++(3M) SDE
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
# Dynamic eta
eta = eta_schedule_cosine_annealing(i, len(sigmas))
h_eta = h * (eta + 1)
x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised
if h_2 is not None:
r0 = h_1 / h
r1 = h_2 / h
d1_0 = (denoised - denoised_1) / r0
d1_1 = (denoised_1 - denoised_2) / r1
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
d2 = (d1_0 - d1_1) / (r0 + r1)
phi_2 = h_eta.neg().expm1() / h_eta + 1
phi_3 = phi_2 / h_eta - 0.5
x = x + phi_2 * d1 - phi_3 * d2
elif h_1 is not None:
r = h_1 / h
d = (denoised - denoised_1) / r
phi_2 = h_eta.neg().expm1() / h_eta + 1
x = x + phi_2 * d
if eta:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise
denoised_1, denoised_2 = denoised, denoised_1
h_1, h_2 = h, h_1
return x
def sample_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback=None, disable=None, eta_max=1.0, eta_min=0.0, s_noise=1., noise_sampler="brownian"):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "highres-pyramid":
noise_sampler = lambda sigma, sigma_next: highres_pyramid_noise_like(x)
case "perlin":
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
case "laplacian":
noise_sampler = lambda sigma, sigma_next: rand_laplacian_like(x)
case _:
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
return sampler_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta_max=eta_max, eta_min=eta_min, s_noise=s_noise, noise_sampler=noise_sampler)
# Add your personal samplers below here, just for formatting purposes ;3
# Add any extra samplers to the following dictionary
@@ -787,6 +808,7 @@ extra_samplers = {
"ttm": sample_ttmcustom,
"lcm_custom_noise": sample_lcmcustom,
"euler_ancestral_dancing": sample_euler_ancestral_dancing,
"dpmpp_3m_sde_dynamic_eta": sample_dpmpp_3m_sde_dynamic_eta,
}
discard_penultimate_sigma_samplers = set((
@@ -794,4 +816,17 @@ discard_penultimate_sigma_samplers = set((
"clyb_4m_sde_momentumized"
))
extra_schedulers = {}
def get_sigmas_simple_exponential(model, steps):
s = model.model_sampling
sigs = []
ss = len(s.sigmas) / steps
for x in range(steps):
sigs += [float(s.sigmas[-(1 + int(x * ss))])]
sigs += [0.0]
sigs = torch.FloatTensor(sigs)
exp = torch.exp(torch.log(torch.linspace(1, 0, steps + 1)))
return sigs * exp
extra_schedulers = {
"simple_exponential": get_sigmas_simple_exponential
}
+48 -3
View File
@@ -7,7 +7,6 @@ import torch
import numpy as np
from tqdm.auto import trange
import random
def pyramid_noise_like(size, dtype, layout, generator, device="cpu", discount=0.8):
b, c, h, w = size
orig_h = h
@@ -189,6 +188,52 @@ class SamplerEULER_ANCESTRAL_DANCING:
sampler = comfy.samplers.ksampler("euler_ancestral_dancing", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise, "leap": leap, "eta_dance": eta_dance})
return (sampler, )
class SamplerDPMPP_3M_SDE_DYN_ETA:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin", "laplacian"], ),
"eta_max": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
"eta_min": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step":0.01}),
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta_max, eta_min, s_noise):
sampler = comfy.samplers.ksampler("dpmpp_3m_sde_dynamic_eta", {"noise_sampler": noise_sampler_type, "eta_max": eta_max, "eta_min": eta_min, "s_noise": s_noise})
return (sampler, )
### Schedulers
from .extra_samplers import get_sigmas_simple_exponential
class SimpleExponentialScheduler:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("SIGMAS",)
CATEGORY = "clybNodes/schedulers"
FUNCTION = "get_sigmas"
def get_sigmas(self, model, steps, denoise):
total_steps = steps
if denoise < 1.0:
total_steps = int(steps/denoise)
sigmas = get_sigmas_simple_exponential(model.model, total_steps).cpu()
sigmas = sigmas[-(steps + 1):]
return (sigmas, )
### KSampler Nodes
from comfy import model_management
import comfy.utils
import comfy.conds
@@ -268,10 +313,10 @@ def mixture_sample(model, model2, noise, positive, positive2, negative, negative
temp_sigmas2 = sigmas2[-2:]
if (i % 2) == 0:
#print(temp_sigmas)
samples = sampler.sample(model_wrap, temp_sigmas, extra_args, callback, noise.to(device) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device), samples if samples is not None else latent_image, denoise_mask, True)
samples = sampler.sample(model_wrap, temp_sigmas, extra_args, callback, noise.to(device) if i == 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device), samples if samples is not None else latent_image, denoise_mask, True)
else:
#print(temp_sigmas)
samples = sampler2.sample(model_wrap2, temp_sigmas2, extra_args2, callback2, noise.to(device2) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device2), samples if samples is not None else latent_image, denoise_mask2, True)
samples = sampler2.sample(model_wrap2, temp_sigmas2, extra_args2, callback2, noise.to(device2) if i == 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device2), samples if samples is not None else latent_image, denoise_mask2, True)
return model.process_latent_out(samples.to(torch.float32))
def sample_mixture(model, model2, noise, cfg, cfg2, sampler, sampler2, sigmas, sigmas2, positive, negative, latent_image, noise_mask=None, callback=None, callback2=None, disable_pbar=False, seed=None):