Add Supreme sampler.
This commit is contained in:
@@ -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
@@ -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((
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user