diff --git a/__init__.py b/__init__.py index 7d72e06..3514fb9 100644 --- a/__init__.py +++ b/__init__.py @@ -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, } diff --git a/extra_samplers.py b/extra_samplers.py index 364f12d..61650b3 100644 --- a/extra_samplers.py +++ b/extra_samplers.py @@ -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(( diff --git a/nodes.py b/nodes.py index 6422196..e0af946 100644 --- a/nodes.py +++ b/nodes.py @@ -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: