From 89a7b35e5e2bd0adc03ff01115fd4c45a4d8d4cd Mon Sep 17 00:00:00 2001 From: Clybius Date: Sun, 5 Nov 2023 14:33:42 -0600 Subject: [PATCH] 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?) --- sampler_mega_modifier.py | 98 ++++++++++++++++++++++++++++++++-------- 1 file changed, 79 insertions(+), 19 deletions(-) diff --git a/sampler_mega_modifier.py b/sampler_mega_modifier.py index 6f64dec..4d46978 100644 --- a/sampler_mega_modifier.py +++ b/sampler_mega_modifier.py @@ -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