From 52eac1b7c847d2727e0ca93ca26d9ffd77029daa Mon Sep 17 00:00:00 2001 From: Clybius Date: Wed, 17 Apr 2024 23:27:33 -0500 Subject: [PATCH] Add MegaCFGGuider and WarmupDecayCFGGuider Add RES step method to Supreme Add spectral noise modulation to Supreme Change reversible dampen to reversible eta on Supreme Remove dyneta temporarily(?) from Supreme Add weight scaling to image/tonal guidance nodes TODO: Update Readme, add start/stop for image guidance, changeable warmup on Supreme (tomorrow) --- __init__.py | 2 + extra_samplers.py | 101 +++++++++++++++++++-- nodes.py | 217 ++++++++++++++++++++++++++++++++++++++++++---- 3 files changed, 297 insertions(+), 23 deletions(-) diff --git a/__init__.py b/__init__.py index 9b02410..a522813 100644 --- a/__init__.py +++ b/__init__.py @@ -13,6 +13,8 @@ NODE_CLASS_MAPPINGS = { "GeometricCFGGuider": nodes.GeometricCFGGuider, "ImageAssistedCFGGuider": nodes.ImageGuidedCFGGuider, "ScaledCFGGuider": nodes.ScaledCFGGuider, + "WarmupDecayCFGGuider": nodes.WarmupDecayCFGGuider, + "MegaCFGGuider": nodes.MegaCFGGuider, ## Samplers "SamplerRES_Momentumized": nodes.SamplerRES_MOMENTUMIZED, "SamplerDPMPP_DualSDE_Momentumized": nodes.SamplerDPMPP_DUALSDE_MOMENTUMIZED, diff --git a/extra_samplers.py b/extra_samplers.py index 5211675..d56d65e 100644 --- a/extra_samplers.py +++ b/extra_samplers.py @@ -795,11 +795,13 @@ def sample_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback= 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 or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args)) +from .other_samplers.refined_exp_solver import _de_second_order + # Default is 2, so only methods with other values are included here. SUPREME_ORDER = { "euler": 1, "dpm_1s": 1, "dpm_3s": 3, "rk4": 4, "reversible_heun_1s": 1, "rkf45": 6, "bogacki_shampine": 3, } @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", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_dampen=1.0): +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", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_eta=1.0): """ Supreme Sampler, Euler steps. Based on no paper, purely interesting thoughts. @@ -821,7 +823,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No noise_modulation: Method of changing the noise based on situations within the sampler modulation_strength: Strength of the modulation using a weighted sum between the modulation and noise sampler's noise. modulation_dims: Choose between (channel) modulation, (height, width) modulation, or (channels, height, width) modulation - reversible_dampen: Power scalar for increasing the strength of the reversible correction dynamically, along with eta and cond modification. + reversible_eta: Power scalar for increasing the strength of the reversible correction dynamically, along with eta and cond modification. """ extra_args = {} if extra_args is None else extra_args @@ -993,6 +995,63 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No scaled_noise = z_k_scaled * intensity + additive_noise * (1 - intensity) return scaled_noise + + def spectral_modulate_noise(z_k, noise, s_noise, sigma_up, intensity, channels, spectral_mod_percentile=5.0): # Modified for soft quantile adjustment using a novel:tm::c::r: method titled linalg. + additive_noise = noise * s_noise * sigma_up + # Convert image to Fourier domain + fourier = torch.fft.fftn(additive_noise, dim=channels) # Apply FFT along Height and Width dimensions + + log_amp = torch.log(torch.sqrt(fourier.real ** 2 + fourier.imag ** 2)) + + quantile_low = torch.quantile( + log_amp.abs().flatten(1), + spectral_mod_percentile * 0.01, + dim = 1 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + quantile_high = torch.quantile( + log_amp.abs().flatten(1), + 1 - (spectral_mod_percentile * 0.01), + dim = 1 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + quantile_max = torch.quantile( + log_amp.abs().flatten(1), + 1, + dim = 1 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + # Decrease high-frequency components + mask_high = log_amp > quantile_high # If we're larger than 95th percentile + + additive_mult_high = torch.where( + mask_high, + 1 - ((log_amp - quantile_high) / (quantile_max - quantile_high)).clamp_(max=0.5), # (1) - (0-1), where 0 is 95th %ile and 1 is 100%ile + torch.tensor(1.0) + ) + + + # Increase low-frequency components + mask_low = log_amp < quantile_low + additive_mult_low = torch.where( + mask_low, + 1 + (1 - (log_amp / quantile_low)).clamp_(max=0.5), # (1) + (0-1), where 0 is 5th %ile and 1 is 0%ile + torch.tensor(1.0) + ) + + mask_mult = ((additive_mult_low * additive_mult_high) ** intensity) + #print(mask_mult) + filtered_fourier = fourier * mask_mult + + # Inverse transform back to spatial domain + inverse_transformed = torch.fft.ifftn(filtered_fourier, dim=channels) # Apply IFFT along Height and Width dimensions + + scaled_noise = inverse_transformed.real.to(additive_noise.device) + + #noise_norm = torch.norm(additive_noise) + #scaled_noise_norm = torch.norm(scaled_noise) + + return scaled_noise# * (noise_norm / scaled_noise_norm) dims = (-3, -2, -1) match modulation_dims: @@ -1021,6 +1080,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No # Renoising iterations z_avg = torch.zeros_like(x) sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + sigma_down_reversible, _ = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=reversible_eta) for k in range(substeps): z_k = x eps_cache = {} @@ -1034,7 +1094,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No step_method_dyn, order, error = dynamic_step_method(step_method, model, prev_x, denoised, prev_denoised, i, k) #step_method, model, prev_x, denoised, prev_denoised, i, k # DynETA - eta = dyneta_fn(orig_eta, error) + #eta = dyneta_fn(orig_eta, error) match step_method_dyn if sigmas[i + 1] != 0 else "euler": case "euler": # 1 model call @@ -1070,6 +1130,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No case "reversible_heun": # 2 model calls sigma_i, sigma_i_plus_1 = sigmas[i], sigma_down dt = sigma_i_plus_1 - sigma_i + dt_reversible = sigma_down_reversible - sigma_i # Calculate the derivative using the model d_i = to_d(z_k, sigma_i, denoised) @@ -1084,11 +1145,12 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, denoised_i_plus_1) # Update the sample using the Reversible Heun formula - z_k = z_k + dt * (d_i + d_i_plus_1) / 2 - dt**2 * (d_i_plus_1 - d_i) / (4 * reversible_dampen) + z_k = z_k + dt * (d_i + d_i_plus_1) / 2 - dt_reversible**2 * (d_i_plus_1 - d_i) / 4 case "reversible_heun_1s": # Experimental 1 model call variant, utilizing previous denoised variables to speed up diffusion. # Reversible Heun-inspired update (first-order) sigma_i, sigma_i_plus_1 = sigmas[i], sigma_down dt = sigma_i_plus_1 - sigma_i + dt_reversible = sigma_down_reversible - sigma_i # Calculate the derivative using the model d_i_old = to_d(prev_x, sigma_i, prev_denoised) if prev_denoised is not None else to_d(prev_x, sigma_i, model(prev_x, sigma_i * s_in, **extra_args)) @@ -1100,7 +1162,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, denoised) # Update the sample using the Reversible Heun formula - z_k = z_k + dt * (d_i_old + d_i_plus_1) / 2 - dt**2 * (d_i_plus_1 - d_i_old) / (2 * reversible_dampen) + z_k = z_k + dt * (d_i_old + d_i_plus_1) / 2 - dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4 case "rkf45": # 6 model calls (expensive) sigma_i, sigma_i_plus_1 = sigmas[i], sigma_down dt = sigma_i_plus_1 - sigma_i @@ -1148,6 +1210,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No case "reversible_bogacki_shampine": sigma_i, sigma_i_plus_1 = sigmas[i], sigma_down dt = sigma_i_plus_1 - sigma_i + dt_reversible = sigma_down_reversible - sigma_i # Calculate the derivative using the model d_i = to_d(z_k, sigma_i, denoised) @@ -1158,7 +1221,7 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No k3 = to_d(z_k + 3 * k1 / 4 + k2 / 4, sigma_i + 3 * dt / 4, model(z_k + 3 * k1 / 4 + k2 / 4, (sigma_i + 3 * dt / 4) * s_in, **extra_args)) * dt # Reversible correction term (inspired by Reversible Heun) - correction = dt**2 * (4 * k3 / 9 - k2 / 3) / (6 * reversible_dampen) + correction = dt_reversible**2 * (k3 - k2) / 6 # Update the sample z_k = z_k + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9 - correction @@ -1183,6 +1246,22 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No z_k = z_k + dt_2 * (d_i + d_i_plus_1) / 2 else: z_k = denoised + case "RES": + lam_next = sigma_down.log().neg() if eta != 0 else sigmas[i + 1].log().neg() + lam = sigmas[i].log().neg() + + h = lam_next - lam + a2_1, b1, b2 = _de_second_order(h=h, c2=0.5, simple_phi_calc=False) + + c2_h = 0.5*h + + x_2 = math.exp(-c2_h)*z_k + a2_1*h*denoised + lam_2 = lam + c2_h + sigma_2 = lam_2.neg().exp() + + denoised2 = model(x_2, sigma_2 * s_in, **extra_args) + + z_k = math.exp(-h)*z_k + h*(b1*denoised + b2*denoised2) z_avg += renoise_weights[k] * z_k if sigmas[i + 1] > 0: # Random noise for variance on ancestral samplers @@ -1196,6 +1275,9 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No case "frequency": noise = noise_sampler(sigmas[i], sigmas[i + 1]) noise_mod = frequency_based_noise(z_k, noise, s_noise, sigma_up, modulation_strength, dims) + case "spectral_signum": + noise = noise_sampler(sigmas[i], sigmas[i + 1]) + noise_mod = spectral_modulate_noise(x, noise, s_noise, sigma_up, modulation_strength, dims) z_k = z_k + noise_mod x = z_avg @@ -1210,6 +1292,9 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No case "frequency": noise = noise_sampler(sigmas[i], sigmas[i + 1]) noise_mod = frequency_based_noise(x, noise, s_noise, sigma_up, modulation_strength, dims) + case "spectral_signum": + noise = noise_sampler(sigmas[i], sigmas[i + 1]) + noise_mod = spectral_modulate_noise(x, noise, s_noise, sigma_up, modulation_strength, dims) x = x + noise_mod @@ -1218,8 +1303,8 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No return x -def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, eta=1.0, step_method="euler", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_dampen=1.0): - return sampler_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), eta=eta, step_method=step_method, substep_method=substep_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_dampen=reversible_dampen) +def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, eta=1.0, step_method="euler", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_eta=1.0): + return sampler_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), eta=eta, step_method=step_method, substep_method=substep_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) # Add your personal samplers below here, just for formatting purposes ;3 diff --git a/nodes.py b/nodes.py index 7fb1fba..1679fa2 100644 --- a/nodes.py +++ b/nodes.py @@ -1,13 +1,19 @@ from .other_samplers.refined_exp_solver import sample_refined_exp_s from .extra_samplers import get_noise_sampler_names, prepare_noise + import comfy.samplers import comfy.sample import comfy.sampler_helpers from comfy.k_diffusion import sampling as k_diffusion_sampling +import node_helpers + import latent_preview import torch +import math from tqdm.auto import trange +import kornia + class SamplerRES_MOMENTUMIZED: @classmethod def INPUT_TYPES(s): @@ -145,9 +151,9 @@ class SamplerDPMPP_3M_SDE_DYN_ETA: class SamplerSUPREME: @classmethod def INPUT_TYPES(s): - SUBSTEP_METHODS=["euler", "dpm_1s", "dpm_2s", "dpm_3s", "bogacki_shampine", "rk4", "rkf45", "reversible_heun", "reversible_heun_1s", "reversible_bogacki_shampine", "trapezoidal"] + SUBSTEP_METHODS=["euler", "dpm_1s", "dpm_2s", "dpm_3s", "bogacki_shampine", "rk4", "rkf45", "reversible_heun", "reversible_heun_1s", "reversible_bogacki_shampine", "trapezoidal", "RES"] STEP_METHODS=SUBSTEP_METHODS+["dynamic", "adaptive_rk"] - NOISE_MODULATION_TYPES=["none", "intensity", "frequency"] + NOISE_MODULATION_TYPES=["none", "intensity", "frequency", "spectral_signum"] return {"required": {"noise_sampler_type": (get_noise_sampler_names(),), "step_method": (STEP_METHODS, ), @@ -162,7 +168,7 @@ class SamplerSUPREME: "noise_modulation": (NOISE_MODULATION_TYPES, {"default": "intensity"}), "modulation_strength": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0, "step":0.01}), "modulation_dims": ("INT", {"default": 3, "min": 1, "max": 3, "step":1}), - "reversible_dampen": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), + "reversible_eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}), } } RETURN_TYPES = ("SAMPLER",) @@ -170,8 +176,8 @@ class SamplerSUPREME: FUNCTION = "get_sampler" - def get_sampler(self, noise_sampler_type, step_method, substep_method, eta, centralization, normalization, edge_enhancement, perphist, substeps, noise_modulation, modulation_strength, modulation_dims, reversible_dampen, s_noise): - sampler = comfy.samplers.ksampler("supreme", {"noise_sampler_type": noise_sampler_type, "step_method": step_method, "eta": eta, "centralization": centralization, "normalization": normalization, "edge_enhancement": edge_enhancement, "perphist": perphist, "substeps": substeps, "substep_method": substep_method, "noise_modulation": noise_modulation, "modulation_strength": modulation_strength, "modulation_dims": modulation_dims, "reversible_dampen": reversible_dampen, "s_noise": s_noise}) + def get_sampler(self, noise_sampler_type, step_method, substep_method, eta, centralization, normalization, edge_enhancement, perphist, substeps, noise_modulation, modulation_strength, modulation_dims, reversible_eta, s_noise): + sampler = comfy.samplers.ksampler("supreme", {"noise_sampler_type": noise_sampler_type, "step_method": step_method, "eta": eta, "centralization": centralization, "normalization": normalization, "edge_enhancement": edge_enhancement, "perphist": perphist, "substeps": substeps, "substep_method": substep_method, "noise_modulation": noise_modulation, "modulation_strength": modulation_strength, "modulation_dims": modulation_dims, "reversible_eta": reversible_eta, "s_noise": s_noise}) return (sampler, ) ### Schedulers @@ -555,11 +561,12 @@ class GeometricCFGGuider: return (guider,) class Guider_ImageGuidedCFG(comfy.samplers.CFGGuider): - def set_cfg(self, model, cfg1, image_cfg, latent_img, img_weighting): + def set_cfg(self, model, cfg1, image_cfg, latent_img, img_weighting, weight_scaling): self.cfg1 = cfg1 self.icfg = image_cfg self.img = latent_img self.img_weighting = img_weighting + self.weight_scaling = weight_scaling self.model = model def set_conds(self, positive, negative): @@ -584,11 +591,13 @@ class Guider_ImageGuidedCFG(comfy.samplers.CFGGuider): case "flat": weight = 1.0 case "linear down": - weight = (self.model.model.model_sampling.timestep(timestep) / 999.0)[:, None, None, None].clone() + weight = (timestep / self.model.model.model_sampling.sigma_max)[:, None, None, None].clone() + case "cosine down": + weight = ((-torch.cos(timestep / self.model.model.model_sampling.sigma_max * math.pi) / 2) + 0.5)[:, None, None, None].clone() cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], self.cfg1, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative_cond) - return cfg + (cfg - res) * self.icfg / self.cfg1 / 10 * weight # Divide by 10 to mimic user-cfg. Do CFG - Res since the image is inverted the other way around. + return cfg + (cfg - res) * self.icfg * (weight**self.weight_scaling) # Divide by 10 to mimic user-cfg. Do CFG - Res since the image is inverted the other way around. class ImageGuidedCFGGuider: @classmethod @@ -598,8 +607,9 @@ class ImageGuidedCFGGuider: "positive": ("CONDITIONING", ), "negative": ("CONDITIONING", ), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "image_cfg": ("FLOAT", {"default": 0.1, "min": -100.0, "max": 100.0, "step":0.1, "round": 0.01}), - "image_weighting": (["flat", "linear down"], ), + "image_cfg": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step":0.01, "round": 0.001}), + "image_weighting": (["flat", "linear down", "cosine down"], ), + "weight_scaling": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100.0, "step":0.01, "round": 0.001}), "latent_image": ("LATENT", ), } } @@ -609,10 +619,10 @@ class ImageGuidedCFGGuider: FUNCTION = "get_guider" CATEGORY = "sampling/custom_sampling/guiders" - def get_guider(self, model, positive, negative, cfg, image_cfg, image_weighting, latent_image): + def get_guider(self, model, positive, negative, cfg, image_cfg, image_weighting, weight_scaling, latent_image): guider = Guider_ImageGuidedCFG(model) guider.set_conds(positive, negative) # Conds - guider.set_cfg(model, cfg, image_cfg, latent_image, image_weighting) # Strengths + guider.set_cfg(model, cfg, image_cfg, latent_image, image_weighting, weight_scaling) # Strengths return (guider,) class Guider_ScaledCFG(comfy.samplers.CFGGuider): @@ -647,7 +657,7 @@ class ScaledCFGGuider: "cond2": ("CONDITIONING", ), "negative": ("CONDITIONING", ), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "cond2_alpha": ("FLOAT", {"default": 1.0, "min": -1.0, "max": 1.0, "step":0.01, "round": 0.01}), + "cond2_alpha": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step":0.01, "round": 0.01}), } } @@ -657,7 +667,184 @@ class ScaledCFGGuider: CATEGORY = "sampling/custom_sampling/guiders" def get_guider(self, model, cond1, cond2, negative, cfg, cond2_alpha): - guider = Guider_GeometricCFG(model) + guider = Guider_ScaledCFG(model) guider.set_conds(cond1, cond2, negative) # Conds guider.set_cfg(cfg, cond2_alpha) # Strengths - return (guider,) \ No newline at end of file + return (guider,) + +class Guider_WarmupDecayCFG(comfy.samplers.CFGGuider): + def set_cfg(self, model, cfg_max, cfg_min, warmup_percent): + self.model = model + self.cfg_max = cfg_max + self.cfg_min = cfg_min + self.warmup_percent = warmup_percent + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + + def predict_noise(self, x, timestep, model_options={}, seed=None): + negative_cond = self.conds.get("negative", None) + positive_cond = self.conds.get("positive", None) + + out = comfy.samplers.calc_cond_batch(self.inner_model, [negative_cond, positive_cond], x, timestep, model_options) # negative, positive2, positive + + sigma_max = self.model.model.model_sampling.sigma_max # 120 + percent_sigma = self.model.model.model_sampling.percent_to_sigma(self.warmup_percent) # 30 + + if timestep > percent_sigma: + decay = (sigma_max - timestep) / (sigma_max - percent_sigma) # (1.0 - (120 - 110) / (120 - 90)) + cfg_scale = 1/2 * (self.cfg_max - self.cfg_min) + cfg_cos = (1 + torch.cos((timestep / sigma_max) * math.pi)) + mod_cfg = cfg_scale * cfg_cos * decay + self.cfg_min + else: + cfg_scale = 1/2 * (self.cfg_max - self.cfg_min) + cfg_cos = (1 + -torch.cos((timestep / percent_sigma) * math.pi)) + mod_cfg = cfg_scale * cfg_cos + self.cfg_min + + cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], mod_cfg, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative_cond) + return cfg + +class WarmupDecayCFGGuider: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"model": ("MODEL",), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "cfg_max": ("FLOAT", {"default": 12.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "cfg_min": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "warmup_percent": ("FLOAT", {"default": 0.5, "min": 0.01, "max": 1.0, "step":0.01, "round": 0.01}), + } + } + + RETURN_TYPES = ("GUIDER",) + + FUNCTION = "get_guider" + CATEGORY = "sampling/custom_sampling/guiders" + + def get_guider(self, model, positive, negative, cfg_max, cfg_min, warmup_percent): + guider = Guider_WarmupDecayCFG(model) + guider.set_conds(positive, negative) # Conds + guider.set_cfg(model, cfg_max, cfg_min, warmup_percent) # Strengths + return (guider,) + +class Guider_MegaCFG(comfy.samplers.CFGGuider): + def set_cfg(self, model, cfg_max, cfg_min, warmup_percent, mean_cfg): + self.model = model + self.cfg_max = cfg_max + self.cfg_min = cfg_min + self.warmup_percent = warmup_percent + self.mean_cfg = mean_cfg + + self.prev_cond = None + self.prev_cfg = None + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + def set_img_cfg(self, image_guidance, image_weighting, weight_scaling, latent_image): + self.image_guidance = image_guidance + self.image_weighting = image_weighting + self.weight_scaling = weight_scaling + self.latent_image = latent_image + + def post_cfg_reference_img(self, args): + model = args["model"] + cond_pred = args["cond_denoised"] + cfg_result = args["denoised"] + sigma = args["sigma"] + + ref = self.latent_image["samples"].to(cfg_result.device) + + if self.image_guidance == 0: + return cfg_result + + norm_out1 = torch.linalg.norm(cond_pred) # Get norm of positive cond + + ref = ref - cond_pred * (cond_pred / norm_out1 * (ref / norm_out1)).sum() # Project positive cond onto image + ref *= torch.linalg.norm(cond_pred) / torch.linalg.norm(ref) # Normalize to cond + ref = self.model.model.model_sampling.calculate_denoised(sigma, ref, cond_pred) + + sigma_max = self.model.model.model_sampling.sigma_max + + weight = 1.0 + match self.image_weighting: + case "linear down": + weight = (sigma / sigma_max)[:, None, None, None].clone() + case "cosine down": + weight = ((-torch.cos((sigma / sigma_max) * math.pi) / 2) + 0.5)[:, None, None, None].clone() + + return cfg_result + (cond_pred - ref) * self.image_guidance * (weight**self.weight_scaling) + + def predict_noise(self, x, timestep, model_options={}, seed=None): + negative_cond = self.conds.get("negative", None) + positive_cond = self.conds.get("positive", None) + + out = comfy.samplers.calc_cond_batch(self.inner_model, [negative_cond, positive_cond], x, timestep, model_options) # negative, positive2, positive + + out0_mean = out[0].mean(dim=(1, 2, 3), keepdim=True) + out1_mean = out[1].mean(dim=(1, 2, 3), keepdim=True) + if self.mean_cfg != 0: + out[0] -= out0_mean + out[1] -= out1_mean + + sigma_max = self.model.model.model_sampling.sigma_max # 120 + percent_sigma = self.model.model.model_sampling.percent_to_sigma(self.warmup_percent) # 30 + + if timestep > percent_sigma: + decay = (sigma_max - timestep) / (sigma_max - percent_sigma) # (1.0 - (120 - 110) / (120 - 90)) + cfg_scale = 1/2 * (self.cfg_max - self.cfg_min) + cfg_cos = (1 + torch.cos((timestep / sigma_max) * math.pi)) + mod_cfg = cfg_scale * cfg_cos * decay + self.cfg_min + else: + cfg_scale = 1/2 * (self.cfg_max - self.cfg_min) + cfg_cos = (1 + -torch.cos((timestep / percent_sigma) * math.pi)) + mod_cfg = cfg_scale * cfg_cos + self.cfg_min + + cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], mod_cfg, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative_cond) + + if self.mean_cfg != 0: + cfg += out0_mean + (out1_mean - out0_mean) * self.mean_cfg + + self.prev_cfg = cfg + self.prev_cond = out[1] + + return cfg + +class MegaCFGGuider: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"model": ("MODEL",), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "cfg_max": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "cfg_min": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + "warmup_percent": ("FLOAT", {"default": 0.5, "min": 0.01, "max": 1.0, "step":0.01, "round": 0.001}), + "mean_cfg": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), + }, + "optional": + { + "image_guidance": ("FLOAT", {"default": 1.0, "min": -1000.0, "max": 1000.0, "step":0.01, "round": 0.001}), + "image_weighting": (["linear down", "cosine down"], ), + "weight_scaling": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100.0, "step":0.01, "round": 0.001}), + "latent_image": ("LATENT", ), + } + } + + RETURN_TYPES = ("GUIDER",) + + FUNCTION = "get_guider" + CATEGORY = "sampling/custom_sampling/guiders" + + def get_guider(self, model, positive, negative, cfg_max, cfg_min, warmup_percent, mean_cfg, + image_guidance, image_weighting, weight_scaling, latent_image = None): + m = model.clone() + guider = Guider_MegaCFG(m) + guider.set_conds(positive, negative) # Conds + guider.set_cfg(m, cfg_max, cfg_min, warmup_percent, mean_cfg) # Strengths + if latent_image != None: + guider.set_img_cfg(image_guidance, image_weighting, weight_scaling, latent_image) + m.set_model_sampler_post_cfg_function(guider.post_cfg_reference_img) + return (guider,)