From efffcbdc3d0cf87265bf06dd052910fd9434c6ec Mon Sep 17 00:00:00 2001 From: Clybius Date: Fri, 15 Nov 2024 11:21:09 -0600 Subject: [PATCH] idk whats even in this but i'll document it and push it later:tm: --- __init__.py | 5 + extra_samplers.py | 638 +++++++++++++++++++++++++++++++++++++++++++--- nodes.py | 241 ++++++++++++++++- 3 files changed, 841 insertions(+), 43 deletions(-) diff --git a/__init__.py b/__init__.py index b6b5ac5..3052e71 100644 --- a/__init__.py +++ b/__init__.py @@ -15,6 +15,7 @@ NODE_CLASS_MAPPINGS = { "ScaledCFGGuider": nodes.ScaledCFGGuider, "WarmupDecayCFGGuider": nodes.WarmupDecayCFGGuider, "MegaCFGGuider": nodes.MegaCFGGuider, + "APGGuider": nodes.APGGuider, ### Noise "ImmiscibleNoise": nodes.ImmiscibleNoise, ## Samplers @@ -29,6 +30,10 @@ NODE_CLASS_MAPPINGS = { "SamplerSENS": nodes.SamplerSENS, "SamplerIPNDM_VAPP": nodes.SamplerIPNDM_VAPP, "SamplerSHIDS": nodes.SamplerSHIDS, + "SamplerDPMPP_2M_SDE_EMA": nodes.SamplerDPMPP_2M_SDE_EMA, + "SamplerBiScope": nodes.SamplerBiScope, + "SamplerEuler_G": nodes.SamplerEuler_G, + "SamplerLeaping_Euler": nodes.SamplerLeaping_Euler, ### Schedulers "SimpleExponentialScheduler": nodes.SimpleExponentialScheduler, "KLOptimalScheduler": nodes.KLOptimalScheduler, diff --git a/extra_samplers.py b/extra_samplers.py index 1309760..61a8a1a 100644 --- a/extra_samplers.py +++ b/extra_samplers.py @@ -50,8 +50,8 @@ def add_schedulers(): # Noise samplers -IMMISCIBLE_NOISE_NAMES=("gaussian_1024", "perlin") -NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian", "immiscible_gaussian", "immiscible_gaussian_maximize", "immiscible_perlin") +IMMISCIBLE_NOISE_NAMES=("gaussian", "perlin") +NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian", "immiscible_gaussian", "immiscible_gaussian_maximize", "immiscible_perlin", "immiscible_perlin_maximize") def get_noise_sampler_names(default=None): if not default: @@ -86,30 +86,35 @@ def check_set_immiscible(x, noise_sampler_type, extra_args): if noise_sampler_type.startswith("immiscible"): match noise_sampler_type: case "immiscible_gaussian": - immiscibility = make_immiscible("gaussian_1024", batching="channel") # FINISH THE REST + immiscibility = make_immiscible("gaussian") # FINISH THE REST extra_args = immiscibility.set_immiscible_extra_args(extra_args) noise_sampler = lambda _sigma, _sigma_next: immiscibility(x) return noise_sampler, extra_args case "immiscible_gaussian_maximize": - immiscibility = make_immiscible("gaussian_1024", maximize=True, batching="channel") # FINISH THE REST + immiscibility = make_immiscible("gaussian", maximize=True) # FINISH THE REST extra_args = immiscibility.set_immiscible_extra_args(extra_args) noise_sampler = lambda _sigma, _sigma_next: immiscibility(x) return noise_sampler, extra_args case "immiscible_perlin": - immiscibility = make_immiscible("perlin", immiscible_latents=8) # FINISH THE REST + immiscibility = make_immiscible("perlin") # FINISH THE REST + extra_args = immiscibility.set_immiscible_extra_args(extra_args) + noise_sampler = lambda _sigma, _sigma_next: immiscibility(x) + return noise_sampler, extra_args + case "immiscible_perlin_maximize": + immiscibility = make_immiscible("perlin", maximize=True) # FINISH THE REST extra_args = immiscibility.set_immiscible_extra_args(extra_args) noise_sampler = lambda _sigma, _sigma_next: immiscibility(x) return noise_sampler, extra_args return None, extra_args class make_immiscible: - def __init__(self, noise_func="gaussian_1024", immiscible_latents=1024, maximize=False, batching="batch"): + def __init__(self, noise_func="gaussian", immiscible_latents=1024, maximize=False): self.noise_func = noise_func self.n_latents = immiscible_latents self.maximize = maximize self.updated_latent = None - self.batching = batching - + + """ def __call__(self, latents): # "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303 # Minimize latent-noise pairs over a batch @@ -166,6 +171,36 @@ class make_immiscible: return noise.view(*xsz[:2], xsz[3], xsz[2]).permute(0, 1, 3, 2) return noise.view(*xsz) + """ + + def __call__(self, latents): + reference_latent = latents + if self.updated_latent != None: + reference_latent = self.updated_latent + + batch_size = latents.shape[0] if self.n_latents is None else self.n_latents + size = [batch_size] + list(latents.shape[1:]) + #noise = torch.randn_like(latents) # [B, C, H, W] + + match self.noise_func: + case "gaussian": + noise = torch.randn(size, dtype=latents.dtype, layout=latents.layout, device=latents.device) + case "perlin": + noise = create_noisy_latents_perlin(torch.randn(size, dtype=latents.dtype, layout=latents.layout, device=latents.device)) + + # Distance calculation (simplified for single process) + distance = torch.linalg.vector_norm( + 0.10 * latents.to(torch.float16).flatten(start_dim=1).unsqueeze(1) - + 0.10 * noise.to(torch.float16).flatten(start_dim=1).unsqueeze(0), + dim=2 + ) # [B, B] + + # Noise Assignment (simplified for single process) + _, col_ind = linear_sum_assignment(distance.cpu().numpy(), maximize=self.maximize) + noise = noise[col_ind].to(latents.device) # Assign the permuted noise + + return noise + def set_immiscible_extra_args(self, extra_args): def immiscible_post_cfg_function(args): self.updated_latent = args["cond_denoised"] @@ -1316,7 +1351,7 @@ def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=Non return sampler_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), eta=eta, step_method=step_method, substep_method=substep_method, warmup_method=warmup_method, centralization=centralization, normalization=normalization, edge_enhancement=edge_enhancement, perphist=perphist, substeps=substeps, noise_modulation=noise_modulation, modulation_strength=modulation_strength, modulation_dims=modulation_dims, reversible_eta=reversible_eta, dyneta=dyneta, reversible_dyneta=reversible_dyneta, enable_free_reverse=enable_free_reverse, free_reverse_eta=free_reverse_eta, free_reverse_dyneta=free_reverse_dyneta) @torch.no_grad() -def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., rsde_eta=1., tsde_eta=1., s_noise=1., noise_sampler=None): +def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., rsde_eta=1., tsde_eta=1., s_noise=1., noise_sampler=None, flow=False): """SDE-Endowed Nimble Sampler. Based off of DPM-Solver++(2M) SDE and DPM-Solver++(3M) SDE. R-SDE for reversible SDE, T-SDE for tertiary SDE.""" if len(sigmas) <= 1: return x @@ -1346,6 +1381,13 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, rsde_eta_h = rsde_eta * h tsde_eta_h = tsde_eta * h + # If/for flow model + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down + renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5 + x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised if old_denoised is not None: @@ -1365,8 +1407,10 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, rphi = tsde_eta_h.neg().expm1() / tsde_eta_h + 1 x = x + rphi * (d + d_2) / 2 - rphi**2 * (d_rev + d_2_rev) / 2 - if eta: + if eta and not flow: x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise + elif eta and flow: + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff old_denoised, old_denoised_2 = denoised, old_denoised h_last, h_last_2 = h, h_last @@ -1376,8 +1420,11 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, def sample_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., rsde_eta=1., tsde_eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None): if len(sigmas) <= 1: return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) - return sampler_sens(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, rsde_eta=rsde_eta, tsde_eta=tsde_eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args)) + return sampler_sens(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, rsde_eta=rsde_eta, tsde_eta=tsde_eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), flow=flow) #From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py #under Apache 2 license @@ -1468,8 +1515,11 @@ def sample_ipndm_vapp(model, x, sigmas, extra_args=None, callback=None, disable= noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) return sampler_ipndm_vapp(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, max_order=max_order, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), pp_guidance=pp_guidance) + +import functools +import operator @torch.no_grad() -def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, order=16): +def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, order=16, eta_order=1., solver_method="weighted_projection", flow=False): """Full ancestral sampling with SHIDS (Stochastic, Historical, Improvised Sampling) 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 @@ -1490,65 +1540,577 @@ def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None old_dt, old_dt_2 = None, None buffer_model_cond = [] - buffer_model_uncond = [] - buffer_model_dt = [] + #buffer_model_uncond = [] + #buffer_model_dt = [] for i in trange(len(sigmas) - 1, disable=disable): denoised = model(x, sigmas[i] * s_in, **extra_args) sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + _, sigma_up_order = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta_order) + + # If/for flow model + downstep_ratio = None + sigma_down_rf = None + alpha_ip1 = None + alpha_down = None + renoise_coeff = None + if flow: + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down_rf = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down_rf + renoise_coeff = (sigmas[i+1]**2 - sigma_down_rf**2*alpha_ip1**2/alpha_down**2)**0.5 + if callback is not None: callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised}) - d_step = to_d(x, sigmas[i], denoised) + d_full = to_d(x, sigmas[i], denoised) d = to_d(x, sigmas[i], temp_uncond[0]) #d_2 = to_d(x, sigmas[i], temp_cond[0]) # Euler method dt = sigma_down - sigmas[i] # Time Difference between now and next step (negative) - x_full = denoised + d_step * sigma_down + x_full = denoised + d_full * sigma_down x_step = denoised + d * sigma_down # Project denoised onto a line between (primarily) x_step (cfgpp), and x_full (normal cfg) - ba = x_step - denoised - ca = x_full - denoised - alpha = (ba * ca) / (ba ** 2 + 1e-8) - x = (1 - alpha)*denoised + alpha*x_step + match solver_method: + case "weighted_projection": + ba = x_step - denoised + ca = x_full - denoised + alpha = (ba * ca) / (ba ** 2 + 1e-8) + x = (1 - alpha)*denoised + alpha*x_step + case "qr_decomposition": + original_shape = x_step.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x_step.reshape(shape_2d) + B = x_full.reshape(shape_2d) + C = denoised.reshape(shape_2d) + Q, _ = torch.qr(A - C) + # Compute the mapping matrix + mapping_matrix = torch.mm(Q.t(), B - C) + mapped_tensor = torch.mm(Q, mapping_matrix) + x = (C + mapped_tensor).reshape(original_shape) + case "svd_lowrank": + original_shape = x_step.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x_step.reshape(shape_2d) + B = x_full.reshape(shape_2d) + C = denoised.reshape(shape_2d) + Ua, Sa, Va = torch.svd_lowrank(A - C, q=6, niter=2) + + A_lowrank = torch.mm(Ua, torch.mm(torch.diag(Sa), Va.t())) + A_diff = (A - C) - A_lowrank + + Qb, _ = torch.qr(B - C) + + A_diff_projected = torch.mm(Qb, torch.mm(Qb.t(), A_diff)) + + x = (B + A_diff_projected).reshape(original_shape) + case "svd": + original_shape = x_step.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x_step.reshape(shape_2d) + B = x_full.reshape(shape_2d) + C = denoised.reshape(shape_2d) + Ua, Sa, Va = torch.linalg.svd(A - C, full_matrices=False, driver="gesvd") + Ub, Sb, Vb = torch.linalg.svd(B - C, full_matrices=False, driver="gesvd")#Sb = torch.linalg.svdvals(B - C, driver="gesvd")# + + A_lowrank = torch.mm(Ub, torch.mm(torch.diag_embed(Sa), Vb)) + A_diff = (A - C) - A_lowrank + + #Qb, _ = torch.qr(B - C) + + #A_diff_projected = torch.mm(Ua, torch.mm(Ua.t(), A_diff)) + + x = (B + A_diff).reshape(original_shape) + + # Create a list of order multipliers + multipliers = [i for i in range(1, len(buffer_model_cond))] + # Normalize so that they're summed up to a total of 1 + total = sum(multipliers) + normalized_multipliers = [m / total for m in multipliers] for iteration in range(len(buffer_model_cond) - 1): - #x = x - (buffer_model_uncond[iteration] - buffer_model_cond[iteration]) / (buffer_model_dt[iteration + 1] / buffer_model_dt[iteration]) - x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up - #denoised = model(x, sigmas[i + 1] * s_in, **extra_args) - ba = x - buffer_model_cond[iteration] - ca = x_step - buffer_model_cond[iteration] - alpha = (ba * ca) / (ba ** 2 + 1e-8) - x = (1 - alpha)*buffer_model_cond[iteration] + alpha*x + if not flow: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up_order * normalized_multipliers[iteration] + elif flow and eta_order: + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta_order + sigma_down_rf = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down_rf + renoise_coeff = (sigmas[i+1]**2 - sigma_down_rf**2*alpha_ip1**2/alpha_down**2)**0.5 + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + match solver_method: + case "weighted_projection": + ba = x - buffer_model_cond[iteration] + ca = x_step - buffer_model_cond[iteration] + alpha = (ba * ca) / (ba ** 2 + 1e-8) + x = (1 - alpha)*buffer_model_cond[iteration] + alpha*x + case "qr_decomposition": + original_shape = x_step.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x.reshape(shape_2d) + B = x_step.reshape(shape_2d) + C = buffer_model_cond[iteration].reshape(shape_2d) + Q, _ = torch.qr(A - C) + # Compute the mapping matrix + mapping_matrix = torch.mm(Q.t(), B - C) + mapped_tensor = torch.mm(Q, mapping_matrix) + x = (C + mapped_tensor).reshape(original_shape) + case "svd_lowrank": + original_shape = x_step.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x.reshape(shape_2d) + B = x_step.reshape(shape_2d) + C = buffer_model_cond[iteration].reshape(shape_2d) + Ua, Sa, Va = torch.svd_lowrank(A - C, q=6, niter=2) + + A_lowrank = torch.mm(Ua, torch.mm(torch.diag(Sa), Va.t())) + A_diff = (A - C) - A_lowrank + + #Qb, _ = torch.qr(B - C) + + #A_diff_projected = torch.mm(Qb, torch.mm(Qb.t(), A_diff)) + + x = (B + A_diff).reshape(original_shape) + case "svd": + original_shape = x.shape + if not original_shape: + shape_2d = (1, 1) + elif len(original_shape) == 4: + shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:])) + else: + shape_2d = (-1, original_shape[-1]) + + A = x.reshape(shape_2d) + B = x_step.reshape(shape_2d) + C = buffer_model_cond[iteration].reshape(shape_2d) + Ua, Sa, Va = torch.linalg.svd(A - C, full_matrices=False, driver="gesvd") + Ub, Sb, Vb = torch.linalg.svd(B - C, full_matrices=False, driver="gesvd")#Sb = torch.linalg.svdvals(B - C, driver="gesvd")# + + A_lowrank = torch.mm(Ub, torch.mm(torch.diag_embed(Sa), Vb)) + A_diff = (A - C) - A_lowrank + + #Qb, _ = torch.qr(B - C) + + #A_diff_projected = torch.mm(Ua, torch.mm(Ua.t(), A_diff)) + + x = (B + A_diff).reshape(original_shape) if len(buffer_model_cond) == max(order - 1, 1): for k in range(order - 2): buffer_model_cond[k] = buffer_model_cond[k+1] - buffer_model_uncond[k] = buffer_model_uncond[k+1] - buffer_model_dt[k] = buffer_model_dt[k+1] + #buffer_model_uncond[k] = buffer_model_uncond[k+1] + #buffer_model_dt[k] = buffer_model_dt[k+1] buffer_model_cond[-1] = denoised.detach() - buffer_model_uncond[-1] = temp_uncond[0].detach() - buffer_model_dt[-1] = dt.detach() + #buffer_model_uncond[-1] = temp_uncond[0].detach() + #buffer_model_dt[-1] = dt.detach() else: buffer_model_cond.append(denoised.detach()) - buffer_model_uncond.append(temp_uncond[0].detach()) - buffer_model_dt.append(dt.detach()) + #buffer_model_uncond.append(temp_uncond[0].detach()) + #buffer_model_dt.append(dt.detach()) #if old_uncond is not None and old_cond is not None and order >= 2: # x = x + (old_cond - old_uncond) / (old_dt / dt) #if old_uncond_2 is not None and old_cond_2 is not None and order >= 3: # x = x + (old_cond_2 - old_uncond_2) / (old_dt_2 / old_dt) / (old_dt / dt) - if sigmas[i + 1] > 0: + if sigmas[i + 1] > 0 and not flow: x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + elif sigmas[i + 1] > 0 and flow: + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff #old_uncond, old_uncond_2 = temp[0], old_uncond #old_cond, old_cond_2 = temp_cond[0], old_cond #old_dt, old_dt_2 = dt, old_dt return x @torch.no_grad() -def sample_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, order=16): +def sample_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, order=16, eta_order=1., solver_method="weighted_projection"): if len(sigmas) <= 1: return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) - return sampler_SHIDS(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), order=order) + return sampler_SHIDS(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), order=order, eta_order=eta_order, solver_method=solver_method, flow=flow) + +@torch.no_grad() +def sampler_dpmpp_2m_sde_ema(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, amp_fac=2., beta1=0.8, beta2=0.95, weight_decay=0.1, centralization=1.0, normalization=1.0, flow=False): + """DPM-Solver++(2M) SDE, with EMA uncond.""" + if len(sigmas) <= 1: + return x + + 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]]) + + old_denoised = None + h_last = None + h = None + + ema = torch.zeros_like(x) + ema_squared = torch.zeros_like(x) + + grad = None + temp_cond = [0] + temp_uncond = [0] + #alpha = torch.linspace(1.0, 0.0, steps=len(sigmas)) ** amp_fac + alpha = [0] + def ema_retrieve_uncond_alpha(args): + temp_cond[0] = args["cond_denoised"] + temp_uncond[0] = args["uncond_denoised"] + alpha[0] = model.inner_model.inner_model.model_sampling.timestep(args["sigma"]) / 999.0 + #alpha[0] = args["sigma"] + return args["denoised"] + + model_options = extra_args.get("model_options", {}).copy() + extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, ema_retrieve_uncond_alpha, disable_cfg1_optimization=True) + + ema = torch.zeros_like(x) + 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++(2M) SDE + t, s = -sigmas[i].log(), -sigmas[i + 1].log() + h = s - t + eta_h = eta * h + + # If/for flow model + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down + renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5 + + grad = denoised + + # Centralization + if centralization != 0: + grad.sub_( + grad.mean(dim=tuple(range(1, grad.dim())), keepdim=True).mul_(centralization) + ) + # Lerp EMA + ema.lerp_(grad, 1. - beta1) + # Normalization + ema.lerp_(ema.div(ema.std(dim=tuple(range(1, grad.dim())), keepdim=True)), weight=normalization) + # Apply EMA onto grad (denoised) + grad.lerp_(ema, beta2) + + if weight_decay != 0: + # Perform stepweight decay + wd_mult = 1 / (1 + weight_decay * (sigmas[i] - sigmas[i + 1])) + grad.mul_(wd_mult) + + ema += (x - temp_uncond[0]) / sigmas[i] * amp_fac * (sigmas[i] - sigmas[i + 1]) + ema -= (x - temp_cond[0]) / sigmas[i] * amp_fac * (sigmas[i] - sigmas[i + 1]) + + x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * grad + + if old_denoised is not None: + r = h_last / h + x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (grad - old_denoised) + + if eta and not flow: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise + elif eta and flow: + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + + old_denoised = denoised + h_last = h + return x + +@torch.no_grad() +def sample_dpmpp_2m_sde_ema(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None, amp_fac=2., beta1=0.8, beta2=0.95, weight_decay=0.1, centralization=1.0, normalization=1.0): + if len(sigmas) <= 1: + return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True + noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) + return sampler_dpmpp_2m_sde_ema(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), amp_fac=amp_fac, beta1=beta1, beta2=beta2, weight_decay=weight_decay, centralization=centralization, normalization=normalization, flow=flow) + +@torch.no_grad() +def sampler_biscope(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, amp_fac=2.0, local_smoothing_fac=4, smoothing_fac=0.75, ema_fac=0.9, flow=False): + """Solving for the Compass model's noise problem using the Compass-like training procedure as an inference sampler.""" + 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]]) + + local_smoothing = [] + smoothing = None + smoothing_diff = None + ema = None + prev_denoised = None + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + + grad = denoised + + if len(local_smoothing) == max(local_smoothing_fac, 1): + for k in range(local_smoothing_fac - 1): + local_smoothing[k] = local_smoothing[k+1] + local_smoothing[-1] = grad.detach() + else: + local_smoothing.append(grad.detach()) + #print(local_smoothing) + local_grad = torch.mean(torch.stack(local_smoothing), dim=0)# if len(local_smoothing) > 1 else grad + + if smoothing is None: + smoothing = local_grad + + smoothing.mul_(smoothing_fac).add_(local_grad, alpha=1 - smoothing_fac) + + diff_grad = local_grad - smoothing + + if smoothing_diff is None: + smoothing_diff = diff_grad + smoothing_diff.mul_(smoothing_fac).add_(diff_grad, alpha=1 - smoothing_fac) + + local_grad.add_(smoothing_diff, alpha=amp_fac) + + if ema is None: + ema = local_grad + ema.mul_(ema_fac).add_(local_grad, alpha=1 - ema_fac) + + sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + + # Flow + downstep_ratio = None + alpha_ip1 = None + alpha_down = None + renoise_coeff = None + if flow: + # If/for flow model + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down + renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5 + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': ema}) + d = to_d(x, sigmas[i], ema) + + # Euler method + dt = sigma_down - sigmas[i] + x = x + d * dt + if sigmas[i + 1] > 0 and not flow: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + elif sigmas[i + 1] > 0 and flow: + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + + prev_denoised = denoised + return x + +@torch.no_grad() +def sample_biscope(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, amp_fac=2.0, local_smoothing_fac=4, smoothing_fac=0.5, ema_fac=0.5): + if len(sigmas) <= 1: + return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True + noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) + return sampler_biscope(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), amp_fac=amp_fac, local_smoothing_fac=local_smoothing_fac, smoothing_fac=smoothing_fac, ema_fac=ema_fac, flow=flow) + +def gaussian_kernel_2d(kernel_size, sigma): + """Generates a 2D Gaussian kernel.""" + k = kernel_size // 2 + x, y = torch.meshgrid(torch.arange(-k, k + 1, dtype=torch.float32), torch.arange(-k, k + 1, dtype=torch.float32)) + gaussian = torch.exp(-(x**2 + y**2) / (2 * sigma**2)) + return gaussian / gaussian.sum() + + +@torch.no_grad() +def sampler_euler_g(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, g_eta=1.0, sigma=5.0, order=2, flow=False): + """Ancestral sampling with Euler method 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 + s_in = x.new_ones([x.shape[0]]) + + neighborhood_size = min(x.shape[-2], x.shape[-1]) * 2 + 1 + padding = neighborhood_size // 2 + kernel = gaussian_kernel_2d(neighborhood_size, sigma).unsqueeze(0).unsqueeze(0).repeat(x.shape[1], 1, 1, 1).to(x.device) + + x_buffer = [] + denoised_buffer = [] + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + if sigmas[i + 1] == 0: + return denoised + sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + # Flow + downstep_ratio = None + alpha_ip1 = None + alpha_down = None + renoise_coeff = None + if flow: + # If/for flow model + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down + renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5 + 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 + + """ + for curr_order in range(1, order): + if sigmas[i + 1] > 0 and not flow: + faux_x = torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigma_up) + elif sigmas[i + 1] > 0 and flow: + #x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + #faux_x = (alpha_ip1/(1 - sigmas[i])) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigmas[i+1]**2 - sigmas[i]**2*alpha_ip1**2/(1 - sigmas[i])**2)**0.5 + faux_x = (alpha_ip1/alpha_down) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + #gauss_x = F.conv2d(faux_x, kernel, padding=padding, groups=x.shape[1]) + faux_denoised = model(faux_x, sigmas[i + 1] * s_in, **extra_args) + faux_d = to_d(faux_x, sigmas[i + 1], faux_denoised) + x = x - faux_d * (sigmas[i + 1] - sigmas[i]) * g_eta / (order - 1) + """ + + if sigmas[i + 1] > 0: + # Create a list of order multipliers + multipliers = [i for i in range(1, len(x_buffer))] + # Normalize so that they're summed up to a total of 1 + total = sum(multipliers) + normalized_multipliers = [m / total for m in multipliers] + + for iteration in range(len(x_buffer) - 1): + if not flow: + faux_x = torch.nn.functional.conv2d(x_buffer[iteration], kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + else: + #x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + #faux_x = (alpha_ip1/(1 - sigmas[i])) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigmas[i+1]**2 - sigmas[i]**2*alpha_ip1**2/(1 - sigmas[i])**2)**0.5 + faux_x = (alpha_ip1/alpha_down) * torch.nn.functional.conv2d(x_buffer[iteration], kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + #gauss_x = F.conv2d(faux_x, kernel, padding=padding, groups=x.shape[1]) + faux_denoised = model(faux_x, sigmas[i + 1] * s_in, **extra_args) + faux_d = to_d(faux_x, sigmas[i + 1], faux_denoised) + x = x - faux_d * (sigmas[i + 1] - sigmas[i]) * g_eta * normalized_multipliers[iteration] + + if len(x_buffer) == max(order - 1, 1): + for k in range(order - 2): + x_buffer[k] = x_buffer[k+1] + denoised_buffer[k] = denoised_buffer[k+1] + x_buffer[-1] = x.detach() + denoised_buffer[-1] = denoised.detach() + else: + x_buffer.append(x.detach()) + denoised_buffer.append(denoised.detach()) + + if flow: + x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + else: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + return x + +@torch.no_grad() +def sample_euler_g(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, g_eta=1.0, sigma=5.0, order=2): + if len(sigmas) <= 1: + return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True + noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) + return sampler_euler_g(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), g_eta=g_eta, sigma=sigma, order=order, flow=flow) + +@torch.no_grad() +def sampler_leaping_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, leap=1, eta=1., s_noise=1., noise_sampler=None, flow=False): + #if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + # return sample_euler_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler) + """Ancestral sampling with Euler method 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 + s_in = x.new_ones([x.shape[0]]) + + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + + do_dance = i < (len(sigmas) - (2 + leap)) + if not do_dance: + leap -= 1 + do_dance = True + + sigma_next = sigmas[i + (1 + leap)] if do_dance else sigmas[i + 1] + sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + + # Flow + downstep_ratio = None + alpha_ip1 = None + alpha_down = None + renoise_coeff = None + if flow: + # If/for flow model + downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta + sigma_down = sigmas[i+1] * downstep_ratio + alpha_ip1 = 1 - sigmas[i+1] + alpha_down = 1 - sigma_down + renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5 + + 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_next - sigmas[i] + x_2 = x + d * dt + + if do_dance: + reverse_denoised = model(x_2, sigma_next * s_in, **extra_args) + _, r_sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + r_d = to_d(x_2, sigma_next, reverse_denoised) + r_dt = sigma_down - sigma_next + x_2 = x + d * dt + r_d * r_dt + if sigmas[i + 1] > 0 and not flow: + x_2 = x_2 + noise_sampler(sigma_next, sigmas[i+1]) * s_noise * sigma_up + elif flow: + x_2 = (alpha_ip1/alpha_down) * x_2 + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff + x = x_2 + + return x + +@torch.no_grad() +def sample_leaping_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, leap=1, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None): + if len(sigmas) <= 1: + return x + flow = False + if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST): + flow = True + noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args) + return sampler_leaping_euler(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, leap=leap, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), flow=flow) # Add your personal samplers below here, just for formatting purposes ;3 @@ -1565,6 +2127,10 @@ extra_samplers = { "sens": sample_sens, "ipndm_vapp": sample_ipndm_vapp, "SHIDS": sample_SHIDS, + "dpmpp_2m_sde_ema": sample_dpmpp_2m_sde_ema, + "biscope": sample_biscope, + "euler_g": sample_euler_g, + "leaping_euler": sample_leaping_euler, } discard_penultimate_sigma_samplers = set(( diff --git a/nodes.py b/nodes.py index fdd791d..5001c7f 100644 --- a/nodes.py +++ b/nodes.py @@ -238,6 +238,8 @@ class SamplerSHIDS: "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), "order": ("INT", {"default": 16, "min": 1, "max": 100, "step":1}), + "eta_order": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "solver_method": (["weighted_projection", "qr_decomposition", "svd_lowrank", "svd"], {"default": "weighted_projection"}), } } RETURN_TYPES = ("SAMPLER",) @@ -245,18 +247,149 @@ class SamplerSHIDS: FUNCTION = "get_sampler" - def get_sampler(self, noise_sampler_type, eta, s_noise, order): - sampler = comfy.samplers.ksampler("SHIDS", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "order": order}) + def get_sampler(self, noise_sampler_type, eta, s_noise, order, eta_order, solver_method): + sampler = comfy.samplers.ksampler("SHIDS", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "order": order, "eta_order": eta_order, "solver_method": solver_method}) + return (sampler, ) + +# EMA DPM++ 2M SDE (Compass Optimizer-like implementation) +class SamplerDPMPP_2M_SDE_EMA: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="brownian"), ), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), + "amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0, "step":0.01}), + "beta1": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}), + "beta2": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 0.999, "step":0.01}), + "weight_decay": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 100.0, "step":0.01}), + "centralization": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step":0.01}), + "normalization": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step":0.01}), + } + } + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, beta1, beta2, weight_decay, centralization, normalization): + sampler = comfy.samplers.ksampler("dpmpp_2m_sde_ema", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "beta1": beta1, "beta2": beta2, "weight_decay": weight_decay, "centralization": centralization, "normalization": normalization}) + return (sampler, ) + +class SamplerEuler_3EMA: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), + "amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}), + "smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}), + "ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}), + "beta": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 0.999, "step":0.01}), + } + } + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, smoothing_fac, ema_fac, beta): + sampler = comfy.samplers.ksampler("euler_3ema", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac, "beta": beta}) + return (sampler, ) + +class SamplerScope: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), + "amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}), + "smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}), + "ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}), + } + } + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, smoothing_fac, ema_fac): + sampler = comfy.samplers.ksampler("scope", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac}) + return (sampler, ) + +class SamplerBiScope: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), + "amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}), + "local_smoothing_fac": ("INT", {"default": 4, "min": 1, "max": 100, "step":1}), + "smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}), + "ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}), + } + } + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, local_smoothing_fac, smoothing_fac, ema_fac): + sampler = comfy.samplers.ksampler("biscope", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "local_smoothing_fac": local_smoothing_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac}) + return (sampler, ) + +class SamplerEuler_G: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "g_eta": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step":0.01}), + "sigma": ("FLOAT", {"default": 1, "min": 0.01, "max": 100.0, "step":0.01}), + "order": ("INT", {"default": 3, "min": 3, "max": 100, "step":1}), + "s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}), + } + } + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, eta, g_eta, sigma, order, s_noise): + sampler = comfy.samplers.ksampler("euler_g", {"noise_sampler_type": noise_sampler_type, "eta": eta, "g_eta": g_eta, "sigma": sigma, "order": order, "s_noise": s_noise}) + return (sampler, ) + +class SamplerLeaping_Euler: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ), + "leap": ("INT", {"default": 1, "min": 1, "max": 8, "step":1}), + "eta": ("FLOAT", {"default": 1.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/samplers" + + FUNCTION = "get_sampler" + + def get_sampler(self, noise_sampler_type, leap, eta, s_noise): + sampler = comfy.samplers.ksampler("leaping_euler", {"noise_sampler_type": noise_sampler_type, "leap": leap, "eta": eta, "s_noise": s_noise}) return (sampler, ) ### Noise class Noise_ImmiscibleNoise: - def __init__(self, noise_type, seed, image_scaling, latent_image): + def __init__(self, noise_type, seed, image_scaling, latent_image, n_latents = 1024): self.noise_type = noise_type self.seed = seed self.image_scaling = image_scaling self.latent_image = latent_image + self.n_latents = n_latents def generate_noise(self, input_latent): latent_image = input_latent["samples"] @@ -264,7 +397,7 @@ class Noise_ImmiscibleNoise: generator = torch.manual_seed(self.seed) if batch_inds is None: gauss = torch.randn_like(latent_image) - noise = make_immiscible(noise_func=self.noise_type)(latent_image if self.latent_image is None else self.latent_image["samples"]) + noise = make_immiscible(noise_func=self.noise_type, immiscible_latents=self.n_latents)(latent_image if self.latent_image is None else self.latent_image["samples"]) noise = gauss * (1.0 - self.image_scaling) + noise * self.image_scaling return noise #return torch.randn(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu") @@ -273,7 +406,7 @@ class Noise_ImmiscibleNoise: noises = [] for i in range(unique_inds[-1]+1): gauss = torch.randn_like(latent_image) - noise = make_immiscible(noise_func=self.noise_type)(latent_image if self.latent_image is None else self.latent_image["samples"]) + noise = make_immiscible(noise_func=self.noise_type, immiscible_latents=self.n_latents)(latent_image if self.latent_image is None else self.latent_image["samples"]) noise = gauss * (1.0 - self.image_scaling) + noise * self.image_scaling if i in unique_inds: noises.append(noise) @@ -288,6 +421,7 @@ class ImmiscibleNoise(DisableNoise): return {"required": { "noise_type": (get_immiscible_noise_sampler_names(), ), + "n_latents": ("INT", {"default": 1024, "min": 1, "max": 16384, "step":1}), "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), }, "optional": @@ -297,8 +431,8 @@ class ImmiscibleNoise(DisableNoise): } } - def get_noise(self, noise_type, noise_seed, image_scaling, latent_image): - return (Noise_ImmiscibleNoise(noise_type, noise_seed, image_scaling, latent_image),) + def get_noise(self, noise_type, n_latents, noise_seed, image_scaling, latent_image): + return (Noise_ImmiscibleNoise(noise_type, noise_seed, image_scaling, latent_image, n_latents),) ### Schedulers from .extra_samplers import get_sigmas_simple_exponential @@ -1057,4 +1191,97 @@ class MegaCFGGuider: if perphist != 0: guider.set_perphist_params(perphist) m.set_model_sampler_post_cfg_function(guider.post_cfg_perphist) + return (guider,) + +class Guider_APG(comfy.samplers.CFGGuider): + class MomentumBuffer: + def __init__(self, momentum: float): + self.momentum = momentum + self.running_average = 0 + def reset(self): + self.running_average = 0 + def update(self, update_value: torch.Tensor): + new_average = self.momentum * self.running_average + self.running_average = update_value + new_average + + def project(self, v0: torch.Tensor, v1: torch.Tensor): + dtype = v0.dtype + v0, v1 = v0.double(), v1.double() + v1 = torch.nn.functional.normalize(v1, dim=[-3, -2, -1]) + v0_parallel = (v0 * v1).sum(dim=[-3, -2, -1], keepdim=True) * v1 + v0_orthogonal = v0 - v0_parallel + return v0_parallel.to(dtype), v0_orthogonal.to(dtype) + + def set_cfg(self, cfg_scale, apg_scale, eta, norm_threshold, momentum_buffer): + self.cfg_scale = cfg_scale + self.apg_scale = apg_scale + self.eta = eta + self.norm_threshold = norm_threshold + self.momentum_buffer = momentum_buffer + + self.curr_timestep = 1. + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + def normalized_guidance(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor, guidance_scale: float, cfg: torch.Tensor = None, momentum_buffer: MomentumBuffer = None, eta: float = 1.0, norm_threshold: float = 0.0): + diff = pred_cond - pred_uncond + + if momentum_buffer is not None: + momentum_buffer.update(diff) + diff = momentum_buffer.running_average + if norm_threshold > 0: + ones = torch.ones_like(diff) + diff_norm = diff.norm(p=2, dim=[-3, -2, -1], keepdim=True) + scale_factor = torch.minimum(ones, norm_threshold / diff_norm) + diff = diff * scale_factor + diff_parallel, diff_orthogonal = self.project(diff, pred_cond) + normalized_update = diff_orthogonal + eta * diff_parallel + if cfg is None: + cfg = pred_cond + pred_guided = cfg + (guidance_scale - 1) * normalized_update + return pred_guided + + def predict_noise(self, x, timestep, model_options={}, seed=None): + negative = self.conds.get("negative", None) + positive_cond = self.conds.get("positive", None) + + # Weird way to workaround momentum buffer sticking from run to run, this should automatically reset it in a majority of cases. + if timestep > (self.curr_timestep - 1e-6): + self.momentum_buffer.reset() + self.curr_timestep = timestep + + out = comfy.samplers.calc_cond_batch(self.inner_model, [negative, positive_cond], x, timestep, model_options) + + cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], self.cfg_scale, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative) + + apg = self.normalized_guidance(out[1], out[0], self.apg_scale, cfg, self.momentum_buffer, self.eta, self.norm_threshold) + + return apg + +class APGGuider: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"model": ("MODEL",), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "apg_scale": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "norm_threshold": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "momentum": ("FLOAT", {"default": -0.5, "min": -100.0, "max": 100.0, "step":0.1, "round": 0.01, "lazy": False}), + } + } + + RETURN_TYPES = ("GUIDER",) + + FUNCTION = "get_guider" + CATEGORY = "sampling/custom_sampling/guiders" + + def get_guider(self, model, positive, negative, cfg_scale, apg_scale, eta, norm_threshold, momentum): + guider = Guider_APG(model) + guider.set_conds(positive, negative) # Conds + momentum_buffer = guider.MomentumBuffer(momentum) + guider.set_cfg(cfg_scale, apg_scale, eta, norm_threshold, momentum_buffer) # Strengths return (guider,) \ No newline at end of file