From 1f161b20f924d26572512046b0fd9ca2702a27dd Mon Sep 17 00:00:00 2001 From: Clybius Date: Mon, 4 Mar 2024 12:32:26 -0600 Subject: [PATCH] Add DPM++ 3M SDE DynETA Partially implement PR #4 & #5 from blepping (oops hold on) --- __init__.py | 3 + extra_samplers.py | 169 ++++++++++++++++++++++++++++------------------ nodes.py | 53 +++++++++++++-- 3 files changed, 154 insertions(+), 71 deletions(-) diff --git a/__init__.py b/__init__.py index d726b5c..7d72e06 100644 --- a/__init__.py +++ b/__init__.py @@ -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'] diff --git a/extra_samplers.py b/extra_samplers.py index a233019..3bb73d1 100644 --- a/extra_samplers.py +++ b/extra_samplers.py @@ -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 = {} \ No newline at end of file +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 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index b9899dc..8ada0cf 100644 --- a/nodes.py +++ b/nodes.py @@ -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 @@ -15,7 +14,7 @@ def pyramid_noise_like(size, dtype, layout, generator, device="cpu", discount=0. noise = torch.zeros(size=size, dtype=dtype, layout=layout, device=device) r = 1 for i in range(5): - r *= 2 # Rather than always going 2x, + r *= 2 # Rather than always going 2x, #w, h = max(1, int(w/(r**i))), max(1, int(h/(r**i))) noise += torch.nn.functional.interpolate((torch.normal(mean=0, std=0.5 ** i, size=(b, c, h * r, w * r), dtype=dtype, layout=layout, generator=generator, device=device)), size=(orig_h, orig_w), mode='nearest-exact') * discount**i #if w>=orig_w*16 or h>=orig_h*16: break @@ -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):