Large update, add spectral_modification, subtract_median combat drift method, update README.md, small changes.

This commit is contained in:
Clybius
2023-09-24 16:50:32 -05:00
parent 303836f04c
commit f7d4dcdda6
2 changed files with 141 additions and 9 deletions
+112 -6
View File
@@ -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