diff --git a/README.md b/README.md index 6864aaa..91e0c7b 100644 --- a/README.md +++ b/README.md @@ -11,10 +11,16 @@ This can lead to more perceptual detail, especially at higher strengths. * Rescale: Scales the CFG by comparing the standard deviation to the existing latent to dynamically lower the CFG. -* Extra Noise: Adds extra noise in the middle of the diffusion process, akin to how sharpness sharpens the noise. +* Extra Noise: Adds extra noise in the middle of the diffusion process to conditioning, and do the inverse operation on unconditioning, if chosen. * Contrast: Adjusts the contrast of the conditioning, can lead to more pop-style results. Essentially functions as a secondary CFG slider for stylization, without changing subject pose and location much, if at all. +* Combat CFG Drift: As we increase CFG, the mean will slightly drift away from 0. This subtracts the mean or median of the latent. Can lead to potentially sharper and higher frequency results, but may result in discoloration. + +* Divisive Norm: Normalizes the latent using avg_pool2d, and can reduce noisy artifacts, due in part to features such as sharpness. + +* Spectral Modulation: Converts the latent to frequencies, and clamps higher frequencies while boosting lower ones, then converts it back to an image latent. This effectively can be treated as a solution to oversaturation or burning as a result of higher CFG values, while not touching values around the median. + ### Tonemapping Methods Explanation: * Reinhard:
Uses the reinhard method of tonemapping (from comfyanonymous' ComfyUI Experiments) to clamp the CFG if the difference is too strong. @@ -26,12 +32,32 @@ This can lead to more perceptual detail, especially at higher strengths. `Closer to 100 percentile == stronger clamping`. Recommended values for testing: tonemap_multiplier of 1, tonemap_percentile of 99.
+* Gated:Clamps the values using torch.quantile, only if above a specific floor value, which is set by `tonemapping_multiplier`. Clamps the noise prediction latent based on the percentile. + + + `Closer to 100 percentile == stronger clamping, lower tonemapping_multiplier == stronger clamping`. Recommended values for testing: tonemap_multiplier of 0.8-1, tonemap_percentile of 99.995.
+* CFG-Mimic:Attempts to mimic a lower or higher CFG based on `tonemapping_multiplier`, and clamps it using `tonemapping_percentile` with torch.quantile. + + + `Closer to 100 percentile == stronger clamping, lower tonemapping_multiplier == stronger clamping`. Recommended values for testing: tonemap_multiplier of 0.33-1.0, tonemap_percentile of 100.
### Contrast Explanation: -Scales the pixel values by the standard deviation, achieving a more contrasty look. In practice, this can effectively act as a secondary CFG slider for stylization. It doesn't modify subject poses much, if at all, which can be great for those looking to get more oomf out of their low-cfg setups.
+Scales the pixel values by the standard deviation, achieving a more contrasty look. In practice, this can effectively act as a secondary CFG slider for stylization. It doesn't modify subject poses much, if at all, which can be great for those looking to get more oomf out of their low-cfg setups. + +Using a negative value will not de-contrast, but instead will use a differing method to do the contrast operation. -33 ought to be near-equivalent to 33 in this case, for example. Feel free to play around and share which you prefer!
+ +### Spectral Modification Explanation: +We boost the low frequencies (low rate of change in the noise), and we lower the high frequencies (high rates of change in the noise). + +Change the low/high frequency range using `spectral_mod_percentile` (default of 5.0, which is the upper and lower 5th percentiles.) + +Increase/Decrease the strength of the adjustment by increasing `spectral_mod_multiplier` + +Beware of percentile values higher than 15, and multiplier values higher than 5. Here be dragons (may cause it to "noise-out", or become full of nonsensical noise, especially earlier in the diffusion process).
+ #### Current Pipeline: ->##### Add extra noise to conditioning -> Sharpen conditioning -> Tonemap conditioning -> Modify contrast of conditioning -> Rescale CFG +>##### Add extra noise to conditioning -> Sharpen conditioning -> Convert to Noise Prediction -> Tonemap Noise Prediction -> Spectral Modification -> Modify contrast of noise prediction -> Rescale CFG -> Divisive Normalization -> Combat CFG Drift #### Why use this over `x` node? Since the `set_model_sampler_cfg_function` hijack in ComfyUI can only utilize a single function, we bundle many latent modification methods into one large function for processing. This is simpler than taking an existing hijack and modifying it, which may be possible, but my (Clybius') lack of Python/PyTorch knowledge leads to this being the optimal method for simplicity. If you know how to do this, feel free to reach out through any means! diff --git a/sampler_mega_modifier.py b/sampler_mega_modifier.py index 53bc1f1..60553ff 100644 --- a/sampler_mega_modifier.py +++ b/sampler_mega_modifier.py @@ -511,6 +511,13 @@ def center_latent_perchannel_with_decorrelate(tensor): # Decorrelates data, slig tensor = flattened.unflatten(2, tensor.shape[2:]) return tensor +def center_latent_median(tensor): + flattened = tensor.flatten(2) + median = flattened.median() + scaled_data = (flattened - median) + scaled_data = scaled_data.unflatten(2, tensor.shape[2:]) + return scaled_data + def divisive_normalization(image_tensor, neighborhood_size, threshold=1e-6): # Compute the local mean and local variance local_mean = F.avg_pool2d(image_tensor, neighborhood_size, stride=1, padding=neighborhood_size // 2, count_include_pad=False) @@ -575,6 +582,86 @@ def get_low_frequency_noise(image: Tensor, threshold: float): return inverse_transformed.real.to(image.device) +def spectral_modulation(image: Tensor, modulation_multiplier: float, spectral_mod_percentile: float): # Reference implementation by Clybius, 2023 :tm::c::r: (jk idc who uses it :3) + # Convert image to Fourier domain + fourier = torch.fft.fft2(image, dim=(-2, -1)) # 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(2), + spectral_mod_percentile * 0.01, + dim = 2 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + quantile_high = torch.quantile( + log_amp.abs().flatten(2), + 1 - (spectral_mod_percentile * 0.01), + dim = 2 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + # Increase low-frequency components + mask_low = ((log_amp < quantile_low).float() + 1).clamp_(max=1.5) # If lower than low 5% quantile, set to 1.5, otherwise 1 + # Decrease high-frequency components + mask_high = ((log_amp < quantile_high).float()).clamp_(min=0.5) # If lower than high 5% quantile, set to 1, otherwise 0.5 + filtered_fourier = fourier * ((mask_low * mask_high) ** modulation_multiplier) # Effectively + + # Inverse transform back to spatial domain + inverse_transformed = torch.fft.ifft2(filtered_fourier, dim=(-2, -1)) # Apply IFFT along Height and Width dimensions + + return inverse_transformed.real.to(image.device) + +def spectral_modulation_soft(image: Tensor, modulation_multiplier: float, spectral_mod_percentile: float): # Modified for soft quantile adjustment using a novel:tm::c::r: method titled linalg. + # Convert image to Fourier domain + fourier = torch.fft.fft2(image, dim=(-2, -1)) # 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(2), + spectral_mod_percentile * 0.01, + dim = 2 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + quantile_high = torch.quantile( + log_amp.abs().flatten(2), + 1 - (spectral_mod_percentile * 0.01), + dim = 2 + ).unsqueeze(-1).unsqueeze(-1).expand(log_amp.shape) + + quantile_max = torch.quantile( + log_amp.abs().flatten(2), + 1, + dim = 2 + ).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) ** modulation_multiplier) + print(mask_mult) + filtered_fourier = fourier * mask_mult + + # Inverse transform back to spatial domain + inverse_transformed = torch.fft.ifft2(filtered_fourier, dim=(-2, -1)) # Apply IFFT along Height and Width dimensions + + return inverse_transformed.real.to(image.device) + class ModelSamplerLatentMegaModifier: @classmethod def INPUT_TYPES(s): @@ -585,7 +672,7 @@ class ModelSamplerLatentMegaModifier: "tonemap_method": (["reinhard", "reinhard_perchannel", "arctan", "quantile", "gated", "cfg-mimic"], ), "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_magnitudes"], ), + "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"], ), @@ -593,6 +680,9 @@ class ModelSamplerLatentMegaModifier: "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}), + "spectral_mod_mode": (["hard_clamp", "soft_clamp"], ), + "spectral_mod_percentile": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 20.0, "step": 0.01}), + "spectral_mod_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 5.0, "step": 0.01}), "affect_uncond": (["None", "Sharpness"], ), }} RETURN_TYPES = ("MODEL",) @@ -600,7 +690,7 @@ class ModelSamplerLatentMegaModifier: 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, 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, spectral_mod_mode, spectral_mod_percentile, spectral_mod_multiplier, affect_uncond): match sharpness_method: case "anisotropic": degrade_func = bilateral_blur @@ -667,10 +757,11 @@ class ModelSamplerLatentMegaModifier: # Sharpness alpha = 1.0 - (timestep / 999.0)[:, None, None, None].clone() # Get alpha multiplier, lower alpha at high sigmas/high noise alpha *= 0.001 * sharpness_multiplier # User-input and weaken the strength so we don't annihilate the latent. - degraded_cond = degrade_func(cond) * alpha + cond * (1.0 - alpha) # Mix the modified latent with the existing latent by the alpha + 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 - noise_pred_degraded = (degraded_cond - uncond) # New noise pred + + noise_pred_degraded = (cond - uncond) # New noise pred # After this point, we use `noise_pred_degraded` instead of just `cond` for the final set of calculations @@ -767,6 +858,18 @@ class ModelSamplerLatentMegaModifier: case _: print("Could not tonemap, for the method was not found.") + # Spectral Modification + if spectral_mod_multiplier > 0: + #alpha = 1. - (timestep / 999.0)[:, None, None, None].clone() # Get alpha multiplier, lower alpha at high sigmas/high noise + #alpha = spectral_mod_multiplier# User-input and weaken the strength so we don't annihilate the latent. + match spectral_mod_mode: + case "hard_clamp": + modulation_func = spectral_modulation + case "soft_clamp": + modulation_func = spectral_modulation_soft + modulation_diff = modulation_func(noise_pred_degraded, spectral_mod_multiplier, spectral_mod_percentile) - noise_pred_degraded + noise_pred_degraded += modulation_diff + if contrast_multiplier > 0: contrast_func = contrast # Contrast, after tonemapping, to ensure user-set contrast is expected to behave similarly across tonemapping settings @@ -784,7 +887,7 @@ class ModelSamplerLatentMegaModifier: x_final = uncond + noise_pred_degraded * cond_scale else: x_cfg = uncond + noise_pred_degraded * cond_scale - ro_pos = torch.std(degraded_cond, dim=(1,2,3), keepdim=True) + ro_pos = torch.std(cond, dim=(1,2,3), keepdim=True) ro_cfg = torch.std(x_cfg, dim=(1,2,3), keepdim=True) x_rescaled = x_cfg * (ro_pos / ro_cfg) @@ -802,9 +905,12 @@ class ModelSamplerLatentMegaModifier: case "subtract": combat_drift_func = center_latent_perchannel alpha = combat_cfg_drift - case "subtract_w_magnitudes": + case "subtract_w_decorrelation": combat_drift_func = center_latent_perchannel_with_decorrelate alpha = combat_cfg_drift + case "subtract_median": + combat_drift_func = center_latent_median + 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