Fix for the ComfyUI sampling update

Added spatial-norm thresholding
Add clamp onto soft spectral modulation for extreme values
Add dynamic cfg modifier (Only halfcosine for now)
Add pyramid noise
Change affect_uncond for sharpness
Stop-gap fix for contrast not operating properly
Stop-gap fix for divisive normalization
Change how combat_cfg_drift works (might change later?)
This commit is contained in:
Clybius
2023-11-05 14:33:42 -06:00
parent 87b3cf3606
commit 89a7b35e5e
+79 -19
View File
@@ -460,6 +460,19 @@ def dyn_thresh_gate(latent: Tensor, centered_magnitudes: Tensor, tonemap_percent
else:
latent = gated_thresholding(tonemap_percentile, floor, latent) # If the magnitudes are higher than the ceiling
return latent # Gated-dynamic thresholding by Birchlabs
def spatial_norm_thresholding(x0, value):
# b c h w
pow_x0 = torch.pow(torch.abs(x0), 2)
s = pow_x0.mean(1, keepdim=True).sqrt().clamp(min=value)
return x0 * (value / s)
def spatial_norm_chw_thresholding(x0, value):
# b c h w
pow_x0 = torch.pow(torch.abs(x0), 2)
s = pow_x0.mean(dim=(1, 2, 3), keepdim=True).sqrt().clamp(min=value)
return x0 * (value / s)
# Contrast function
def contrast(x: Tensor):
@@ -653,8 +666,8 @@ def spectral_modulation_soft(image: Tensor, modulation_multiplier: float, spectr
torch.tensor(1.0)
)
mask_mult = ((additive_mult_low * additive_mult_high) ** modulation_multiplier)
mask_mult = ((additive_mult_low * additive_mult_high) ** modulation_multiplier).clamp_(min=0.05, max=20)
#print(mask_mult)
filtered_fourier = fourier * mask_mult
# Inverse transform back to spatial domain
@@ -662,6 +675,40 @@ def spectral_modulation_soft(image: Tensor, modulation_multiplier: float, spectr
return inverse_transformed.real.to(image.device)
import random
def pyramid_noise_like(x, discount=0.9):
b, c, w, h = x.shape # EDIT: w and h get over-written, rename for a different variant!
u = torch.nn.Upsample(size=(w, h), mode='nearest-exact')
noise = torch.randn_like(x)
for i in range(10):
r = random.random()*2+2 # Rather than always going 2x,
w, h = max(1, int(w/(r**i))), max(1, int(h/(r**i)))
noise += u(torch.randn(b, c, w, h).to(x)) * discount**i
if w==1 or h==1: break # Lowest resolution is 1x1
return noise/noise.std() # Scaled back to roughly unit variance
import math
def dyn_cfg_modifier(conditioning, unconditioning, method, cond_scale, time_mult):
match method:
case "dyncfg-halfcosine":
noise_pred = conditioning - unconditioning
#noise_pred = noise_pred * (cond_scale * time)
noise_pred_magnitude = (torch.linalg.vector_norm(noise_pred, dim=(1)) + 0.0000000001)[:,None]
time = time_mult.item()
time_factor = -(math.cos(0.5 * time * math.pi) / 2) + 1
noise_pred_timescaled_magnitude = (torch.linalg.vector_norm(noise_pred * time_factor, dim=(1)) + 0.0000000001)[:,None]
#distance = torch.dist(noise_pred_magnitude, noise_pred_timescaled_magnitude, 2)
#print(distance)
#cos_pred = (conditioning * distance) - unconditioning
#cos_pred_vector_magnitude = (torch.linalg.vector_norm(noise_pred, dim=(1)) + 0.0000000001)[:,None]
noise_pred /= noise_pred_magnitude
noise_pred *= noise_pred_timescaled_magnitude
return noise_pred
class ModelSamplerLatentMegaModifier:
@classmethod
def INPUT_TYPES(s):
@@ -669,28 +716,30 @@ class ModelSamplerLatentMegaModifier:
"sharpness_multiplier": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.1}),
"sharpness_method": (["anisotropic", "gaussian"], ),
"tonemap_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.01}),
"tonemap_method": (["reinhard", "reinhard_perchannel", "arctan", "quantile", "gated", "cfg-mimic"], ),
"tonemap_method": (["reinhard", "reinhard_perchannel", "arctan", "quantile", "gated", "cfg-mimic", "spatial-norm"], ),
"tonemap_percentile": ("FLOAT", {"default": 100.0, "min": 0.0, "max": 100.0, "step": 0.005}),
"contrast_multiplier": ("FLOAT", {"default": 0.0, "min": -100.0, "max": 100.0, "step": 0.1}),
"combat_method": (["subtract", "subtract_w_decorrelation", "subtract_median"], ),
"combat_cfg_drift": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"rescale_cfg_phi": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"extra_noise_type": (["gaussian", "uniform", "perlin", "pink", "green"], ),
"extra_noise_type": (["gaussian", "uniform", "perlin", "pink", "green", "pyramid"], ),
"extra_noise_method": (["add", "add_scaled", "speckle"], ),
"extra_noise_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}),
"extra_noise_lowpass": ("INT", {"default": 100, "min": 0, "max": 1000, "step": 1}),
"divisive_norm_size": ("INT", {"default": 0, "min": 0, "max": 31, "step": 1}),
"divisive_norm_size": ("INT", {"default": 127, "min": 1, "max": 255, "step": 1}),
"divisive_norm_multiplier": ("FLOAT", {"default": 0, "min": 0, "max": 1, "step": 0.01}),
"spectral_mod_mode": (["hard_clamp", "soft_clamp"], ),
"spectral_mod_percentile": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 50.0, "step": 0.01}),
"spectral_mod_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 15.0, "step": 0.01}),
"affect_uncond": (["None", "Sharpness"], ),
"dyn_cfg_augmentation": (["None", "dyncfg-halfcosine"], ),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "mega_modify"
CATEGORY = "clybNodes"
def mega_modify(self, model, sharpness_multiplier, sharpness_method, tonemap_multiplier, tonemap_method, tonemap_percentile, contrast_multiplier, combat_method, combat_cfg_drift, rescale_cfg_phi, extra_noise_type, extra_noise_method, extra_noise_multiplier, extra_noise_lowpass, divisive_norm_size, spectral_mod_mode, spectral_mod_percentile, spectral_mod_multiplier, affect_uncond):
def mega_modify(self, model, sharpness_multiplier, sharpness_method, tonemap_multiplier, tonemap_method, tonemap_percentile, contrast_multiplier, combat_method, combat_cfg_drift, rescale_cfg_phi, extra_noise_type, extra_noise_method, extra_noise_multiplier, extra_noise_lowpass, divisive_norm_size, divisive_norm_multiplier, spectral_mod_mode, spectral_mod_percentile, spectral_mod_multiplier, affect_uncond, dyn_cfg_augmentation):
match sharpness_method:
case "anisotropic":
degrade_func = bilateral_blur
@@ -703,7 +752,8 @@ class ModelSamplerLatentMegaModifier:
cond = args["cond"]
uncond = args["uncond"]
cond_scale = args["cond_scale"]
timestep = args["timestep"]
timestep = model.model.model_sampling.timestep(args["timestep"])
#print(model.model.model_sampling.timestep(timestep))
noise_pred = (cond - uncond)
# Extra noise
@@ -735,6 +785,8 @@ class ModelSamplerLatentMegaModifier:
std = torch.std(extra_noise)
extra_noise.sub_(mean).div_(std)
case "pyramid":
extra_noise = pyramid_noise_like(cond)
if extra_noise_lowpass > 0:
extra_noise = get_low_frequency_noise(extra_noise, extra_noise_lowpass)
@@ -759,9 +811,10 @@ class ModelSamplerLatentMegaModifier:
alpha *= 0.001 * sharpness_multiplier # User-input and weaken the strength so we don't annihilate the latent.
cond = degrade_func(cond) * alpha + cond * (1.0 - alpha) # Mix the modified latent with the existing latent by the alpha
if affect_uncond == "Sharpness":
uncond = uncond + (uncond - degrade_func(uncond)) * alpha
uncond = degrade_func(uncond) * alpha + uncond * (1.0 - alpha)
noise_pred_degraded = (cond - uncond) # New noise pred
time_mult = 1.0 - (timestep / 999.0)[:, None, None, None].clone()
noise_pred_degraded = (cond - uncond) if dyn_cfg_augmentation == "None" else dyn_cfg_modifier(cond, uncond, dyn_cfg_augmentation, cond_scale, time_mult) # New noise pred
# After this point, we use `noise_pred_degraded` instead of just `cond` for the final set of calculations
@@ -855,6 +908,10 @@ class ModelSamplerLatentMegaModifier:
pred_renorm = pred_normalized * mimic_max
pred_uncentered = pred_renorm + mimic_means # Personal choice to re-mean from the mimic here... should be latent_means.
noise_pred_degraded = pred_uncentered.unflatten(2, noise_pred_degraded.shape[2:])
case "spatial-norm":
time = (1.0 - (timestep / 999.0)[:, None, None, None].clone().item())
time = -(math.cos(time * math.pi) / (3)) + (2/3)
noise_pred_degraded = spatial_norm_chw_thresholding(noise_pred_degraded, (tonemap_multiplier / cond_scale) * time)
case _:
print("Could not tonemap, for the method was not found.")
@@ -875,12 +932,12 @@ class ModelSamplerLatentMegaModifier:
# Contrast, after tonemapping, to ensure user-set contrast is expected to behave similarly across tonemapping settings
alpha = 1.0 - (timestep / 999.0)[:, None, None, None].clone()
alpha *= 0.001 * contrast_multiplier
noise_pred_degraded = contrast_func(noise_pred_degraded) * alpha + noise_pred_degraded * (1.0 - alpha)
noise_pred_degraded = (contrast_func(noise_pred_degraded + args["input"]) * alpha + (noise_pred_degraded + args["input"]) * (1.0 - alpha)) - args["input"] # Temporary fix for contrast is to add the input? Maybe? It just doesn't work like before...
if contrast_multiplier < 0:
contrast_func = contrast_with_mean # Unsure if good/bad, buuut its a nice alternative to combatting cfg drift directly
alpha = 1.0 - (timestep / 999.0)[:, None, None, None].clone()
alpha *= 0.001 * -contrast_multiplier # Since we're less than 0, we kinda wanna do the function properly!
noise_pred_degraded = contrast_func(noise_pred_degraded) * alpha + noise_pred_degraded * (1.0 - alpha)
noise_pred_degraded = (contrast_func(noise_pred_degraded + args["input"]) * alpha + (noise_pred_degraded + args["input"]) * (1.0 - alpha)) - args["input"]
# Rescale CFG
if rescale_cfg_phi == 0:
@@ -893,24 +950,27 @@ class ModelSamplerLatentMegaModifier:
x_rescaled = x_cfg * (ro_pos / ro_cfg)
x_final = rescale_cfg_phi * x_rescaled + (1.0 - rescale_cfg_phi) * x_cfg
if divisive_norm_size > 0:
if divisive_norm_multiplier > 0:
alpha = 1. - (timestep / 999.0)[:, None, None, None].clone()
alpha ** 0.1 # Alpha might as well be 1, but we want to protect the beginning steps (?).
high_noise = divisive_normalization(x_final, (divisive_norm_size * 2) + 1)
#alpha *= 0.1 # Alpha might as well be 1, but we want to protect the beginning steps (?).
alpha *= divisive_norm_multiplier
high_noise = args["input"] - divisive_normalization(args["input"] - x_final, (divisive_norm_size * 2) + 1)
x_final = high_noise * alpha + x_final * (1.0 - alpha)
is_early_step = (timestep / 999.0)[:, None, None, None].clone() > 0.8
if combat_cfg_drift > 0 and not is_early_step:
if combat_cfg_drift > 0:
alpha = (1. - (timestep / 999.0)[:, None, None, None].clone()) * 2
alpha ** 0.025 # Alpha might as well be 1, but we want to protect the beginning steps (?).
alpha = alpha.clamp_(max=1)
match combat_method:
case "subtract":
combat_drift_func = center_latent_perchannel
alpha = combat_cfg_drift
alpha *= combat_cfg_drift
case "subtract_w_decorrelation":
combat_drift_func = center_latent_perchannel_with_decorrelate
alpha = combat_cfg_drift
alpha *= combat_cfg_drift
case "subtract_median":
combat_drift_func = center_latent_median
alpha = combat_cfg_drift
alpha *= combat_cfg_drift
x_final = combat_drift_func(x_final) * alpha + x_final * (1.0 - alpha) # Mix the modified latent with the existing latent by the alpha
return x_final # General formula for CFG. uncond + (cond - uncond) * cond_scale