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)
This commit is contained in:
Clybius
2024-04-17 23:27:33 -05:00
parent ed17033a5a
commit 52eac1b7c8
3 changed files with 297 additions and 23 deletions
+2
View File
@@ -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,
+93 -8
View File
@@ -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
+202 -15
View File
@@ -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,)
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,)