Add Supreme sampler.

This commit is contained in:
Clybius
2024-03-20 19:00:47 -05:00
parent b609db0df7
commit 84c831ef24
3 changed files with 203 additions and 1 deletions
+1
View File
@@ -15,6 +15,7 @@ NODE_CLASS_MAPPINGS = {
"SamplerLCMCustom": nodes.SamplerLCMCustom,
"SamplerEulerAncestralDancing_Experimental": nodes.SamplerEULER_ANCESTRAL_DANCING,
"SamplerDPMPP_3M_SDE_DynETA": nodes.SamplerDPMPP_3M_SDE_DYN_ETA,
"SamplerSupreme": nodes.SamplerSUPREME,
### Schedulers
"SimpleExponentialScheduler": nodes.SimpleExponentialScheduler,
}
+179 -1
View File
@@ -7,7 +7,7 @@ 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
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, PIDStepSizeController, get_ancestral_step, to_d, default_noise_sampler, DPMSolver
# The following function adds the samplers during initialization, in __init__.py
def add_samplers():
@@ -798,6 +798,183 @@ def sample_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback=
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)
@torch.no_grad()
def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler=None, eta=1.0, step_method="euler", centralization=0.02, normalization=0.01, edge_enhancement=0.5, perphist=-0.15):
"""
Supreme Sampler, Euler steps. Based on no paper, purely interesting thoughts.
Args:
model: Denoising model call.
x: The initial noisy sample.
sigmas: The noise schedule.
extra_args: Additional arguments for the model.
callback: A callback function for monitoring the sampling process.
disable: Whether to disable the progress bar.
s_noise: The noise scale factor.
noise_sampler: A custom noise sampler function.
eta: Ancestral-ness.
centralization: Subtracts mean from the denoised latent, reduces edge enhancement when between 0-1.
normalization: Divides the denoised latent by the standard deviation.
edge_enhancement: Multiplies the edges by the mean using a laplacian kernel
perphist: Adds previous denoised variable to the current denoised using perpendicular vector projection
"""
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]])
# Centralization
def centralize(denoised_sample, centralization):
for b in range(len(denoised_sample)):
for c in range(len(denoised_sample[b])):
channel = denoised_sample[b][c]
denoised_sample[b][c] -= channel.mean() * centralization
return denoised_sample
# Normalization
def normalize(denoised_sample, normalization):
for b in range(len(denoised_sample)):
for c in range(len(denoised_sample[b])):
channel = denoised_sample[b][c]
denoised_sample[b][c] += ((denoised_sample[b][c] / channel.std()) - denoised_sample[b][c]) * normalization
return denoised_sample
# Perp-hist
def perpadd(denoised_tensor, old_denoised_tensor, x, alpha):
a_diff = x - (denoised_tensor - x)
b_diff = x - (old_denoised_tensor - x)
a_ortho = a_diff * (a_diff / torch.linalg.norm(a_diff) * (b_diff / torch.linalg.norm(a_diff))).sum()
b_perp = b_diff - a_ortho
res = denoised_tensor + alpha * b_perp
return res
# DynETA
def eta_schedule_cosine_annealing(i, n, eta_max=eta, eta_min=0.0):
"""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
def f(x, sigma):
"""Function representing the Karras ODE derivative."""
denoised = model(x, sigma * s_in, **extra_args)
return to_d(x, sigma, denoised)
old_denoised = None
for i in trange(len(sigmas) - 1, disable=disable):
# DynETA
eta = eta_schedule_cosine_annealing(i, len(sigmas))
eps_cache = {}
dpm_solver = DPMSolver(model, extra_args)
denoised = model(x, sigmas[i] * s_in, **extra_args)
if centralization != 0:
denoised = centralize(denoised, centralization)
if normalization != 0:
denoised = normalize(denoised, normalization)
if old_denoised != None and perphist != 0:
denoised = perpadd(denoised, old_denoised, x, perphist)
if old_denoised != None and edge_enhancement != 0:
lap_kern = torch.tensor([[0, -1, 0], [-1, 4, -1], [0, -1, 0]], device=denoised.device, dtype=denoised.dtype).repeat(denoised.shape[1], 1, 1, 1)
denoised = denoised + torch.conv2d(denoised, lap_kern, groups=denoised.shape[1], padding=1) * denoised.mean(dim=(1, 2, 3), keepdim=True) * sigmas[i] * edge_enhancement
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
eps = (x - denoised) / sigmas[i]
eps_cache = {'eps': eps}
match step_method:
case "euler":
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
d = to_d(x, sigmas[i], denoised)
dt = sigma_down - sigmas[i]
x = x + d * dt
if sigmas[i + 1] > 0:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
case "dpm_1s": # DPM Family
if callback is not None:
dpm_solver.info_callback = lambda info: callback({'sigma': dpm_solver.sigma(info['t']), 'sigma_hat': dpm_solver.sigma(info['t_up']), **info})
x, eps_cache = dpm_solver.dpm_solver_1_step(x, dpm_solver.t(sigmas[i]), dpm_solver.t(sigmas[i + 1]), eps_cache=eps_cache)
case "dpm_2s":
if callback is not None:
dpm_solver.info_callback = lambda info: callback({'sigma': dpm_solver.sigma(info['t']), 'sigma_hat': dpm_solver.sigma(info['t_up']), **info})
x, eps_cache = dpm_solver.dpm_solver_2_step(x, dpm_solver.t(sigmas[i]), dpm_solver.t(sigmas[i + 1]), eps_cache=eps_cache)
case "dpm_3s":
if callback is not None:
dpm_solver.info_callback = lambda info: callback({'sigma': dpm_solver.sigma(info['t']), 'sigma_hat': dpm_solver.sigma(info['t_up']), **info})
x, eps_cache = dpm_solver.dpm_solver_3_step(x, dpm_solver.t(sigmas[i]), dpm_solver.t(sigmas[i + 1]), eps_cache=eps_cache)
case "rk4": # Fourth-order Runge-Kutta method
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
# Calculate the derivative using the model
d = to_d(x, sigmas[i], denoised)
dt = sigma_down - sigmas[i]
# Runge-Kutta steps
k1 = d * dt
k2 = to_d(x + k1 / 2, sigmas[i] + dt / 2, model(x + k1 / 2, (sigmas[i] + dt / 2) * s_in, **extra_args)) * dt
k3 = to_d(x + k2 / 2, sigmas[i] + dt / 2, model(x + k2 / 2, (sigmas[i] + dt / 2) * s_in, **extra_args)) * dt
k4 = to_d(x + k3, sigmas[i] + dt, model(x + k3, (sigmas[i] + dt) * s_in, **extra_args)) * dt
# Update the sample
x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6
if sigmas[i + 1] > 0:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
case "trapezoidal":
if sigmas[i + 1] > 0:
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
dt = sigmas[i + 1] - sigmas[i]
# Calculate the derivative using the model
d_i = to_d(x, sigmas[i], denoised)
# Predict the sample at the next sigma using Euler step
x_pred = x + d_i * dt
# Denoised sample at the next sigma
denoised_i_plus_1 = model(x_pred, sigmas[i + 1] * s_in, **extra_args)
# Calculate the derivative at the next sigma
d_i_plus_1 = to_d(x_pred, sigmas[i + 1], denoised_i_plus_1)
#if callback is not None:
# callback({'x': x, 'i': i, 'sigma': sigmas[i + 1], 'denoised': denoised_i_plus_1})
dt_2 = sigma_down - sigmas[i]
# Update the sample using the Trapezoidal rule
x = x + dt_2 * (d_i + d_i_plus_1) / 2
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
else:
x = denoised
old_denoised = denoised
return x
def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler="gaussian", eta=1.0, step_method="euler", centralization=0.02, normalization=0.01, edge_enhancement=0.5, perphist=-0.15):
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_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler, eta=eta, step_method=step_method, centralization=centralization, normalization=normalization, edge_enhancement=edge_enhancement, perphist=perphist)
# Add your personal samplers below here, just for formatting purposes ;3
# Add any extra samplers to the following dictionary
@@ -809,6 +986,7 @@ extra_samplers = {
"lcm_custom_noise": sample_lcmcustom,
"euler_ancestral_dancing": sample_euler_ancestral_dancing,
"dpmpp_3m_sde_dynamic_eta": sample_dpmpp_3m_sde_dynamic_eta,
"supreme": sample_supreme,
}
discard_penultimate_sigma_samplers = set((
+23
View File
@@ -207,6 +207,29 @@ class SamplerDPMPP_3M_SDE_DYN_ETA:
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, )
class SamplerSUPREME:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin", "laplacian"], ),
"step_method": (["euler", "dpm_1s", "dpm_2s", "dpm_3s", "rk4", "trapezoidal"], ),
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
"centralization": ("FLOAT", {"default": 0.02, "min": -1.0, "max": 1.0, "step":0.01}),
"normalization": ("FLOAT", {"default": 0.01, "min": -1.0, "max": 1.0, "step":0.01}),
"edge_enhancement": ("FLOAT", {"default": 0.5, "min": -100.0, "max": 100.0, "step":0.01}),
"perphist": ("FLOAT", {"default": -0.15, "min": -5.0, "max": 5.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, step_method, eta, centralization, normalization, edge_enhancement, perphist, s_noise):
sampler = comfy.samplers.ksampler("supreme", {"noise_sampler": noise_sampler_type, "step_method": step_method, "eta": eta, "centralization": centralization, "normalization": normalization, "edge_enhancement": edge_enhancement, "perphist": perphist, "s_noise": s_noise})
return (sampler, )
### Schedulers
from .extra_samplers import get_sigmas_simple_exponential
class SimpleExponentialScheduler: