Added two new tonemapping methods (gated, cfg-mimic)
Added affect_uncond (for sharpness) Made extra noise affect uncond Added low-pass for extra_noise (May go away?) Added combat cfg drift (mean centering) Added a new contrast method (use negative values for usage) Added divisive normalization, which can help with extreme added values (like sharpness)
This commit is contained in:
+229
-10
@@ -440,9 +440,37 @@ def train_difference(a: Tensor, b: Tensor, c: Tensor) -> Tensor:
|
||||
new_diff = scale * torch.abs(diff_AB)
|
||||
return new_diff
|
||||
|
||||
def gated_thresholding(percentile: float, floor: float, t: Tensor) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
percentile: float between 0.0 and 1.0. for example 0.995 would subject only the top 0.5%ile to clamping.
|
||||
t: [b, c, v] tensor in pixel or latent space (where v is the result of flattening w and h)
|
||||
"""
|
||||
a = t.abs() # Magnitudes
|
||||
q = torch.quantile(a, percentile, dim=2) # Get clamp value via top % of magnitudes
|
||||
q.clamp_(min=floor)
|
||||
q = q.unsqueeze(2).expand(*t.shape)
|
||||
t = t.clamp(-q, q) # Clamp latent with magnitude value
|
||||
t = t / q
|
||||
return t
|
||||
|
||||
def dyn_thresh_gate(latent: Tensor, centered_magnitudes: Tensor, tonemap_percentile: float, floor: float, ceil: float):
|
||||
if centered_magnitudes.lt(torch.tensor(ceil, device=centered_magnitudes.device)).all().item(): # If the magnitudes are less than the ceiling
|
||||
return latent # Return the unmodified centered latent
|
||||
else:
|
||||
latent = gated_thresholding(tonemap_percentile, floor, latent) # If the magnitudes are higher than the ceiling
|
||||
return latent # Gated-dynamic thresholding by Birchlabs
|
||||
# Contrast function
|
||||
|
||||
def contrast(x: Tensor):
|
||||
# Calculate the mean and standard deviation of the pixel values
|
||||
#mean = x.mean(dim=(1,2,3), keepdim=True)
|
||||
stddev = x.std(dim=(1,2,3), keepdim=True)
|
||||
# Scale the pixel values by the standard deviation
|
||||
scaled_pixels = (x) / stddev
|
||||
return scaled_pixels
|
||||
|
||||
def contrast_with_mean(x: Tensor):
|
||||
# Calculate the mean and standard deviation of the pixel values
|
||||
mean = x.mean(dim=(1,2,3), keepdim=True)
|
||||
stddev = x.std(dim=(1,2,3), keepdim=True)
|
||||
@@ -450,6 +478,103 @@ def contrast(x: Tensor):
|
||||
scaled_pixels = (x - mean) / stddev
|
||||
return scaled_pixels
|
||||
|
||||
def center_latent(tensor): #https://birchlabs.co.uk/machine-learning#combating-mean-drift-in-cfg
|
||||
"""Centers on 0 to combat CFG drift."""
|
||||
tensor = tensor - tensor.mean(dim=(-2, -1)).unsqueeze(-1).unsqueeze(-1).expand(tensor.shape)
|
||||
return tensor
|
||||
|
||||
def center_latent_perchannel(tensor): # Does nothing different than above
|
||||
"""Centers on 0 to combat CFG drift."""
|
||||
flattened = tensor.flatten(2)
|
||||
flattened = flattened - flattened.mean(dim=(2)).unsqueeze(2).expand(flattened.shape)
|
||||
tensor = flattened.unflatten(2, tensor.shape[2:])
|
||||
return tensor
|
||||
|
||||
def center_latent_perchannel_with_magnitudes(tensor): # Does nothing different than above
|
||||
"""Centers on 0 to combat CFG drift."""
|
||||
flattened = tensor.flatten(2)
|
||||
flattened_magnitude = (torch.linalg.vector_norm(flattened, dim=(2), keepdim=True) + 0.0000000001)
|
||||
flattened /= flattened_magnitude
|
||||
flattened = flattened - flattened.mean(dim=(2)).unsqueeze(2).expand(flattened.shape)
|
||||
flattened *= flattened_magnitude
|
||||
tensor = flattened.unflatten(2, tensor.shape[2:])
|
||||
return tensor
|
||||
|
||||
def center_latent_perchannel_with_decorrelate(tensor): # Decorrelates data, slight change, test and play with it.
|
||||
"""Centers on 0 to combat CFG drift, preprocesses the latent with decorrelation"""
|
||||
tensor = decorrelate_data(tensor)
|
||||
flattened = tensor.flatten(2)
|
||||
flattened_magnitude = (torch.linalg.vector_norm(flattened, dim=(2), keepdim=True) + 0.0000000001)
|
||||
flattened /= flattened_magnitude
|
||||
flattened = flattened - flattened.mean(dim=(2)).unsqueeze(2).expand(flattened.shape)
|
||||
flattened *= flattened_magnitude
|
||||
tensor = flattened.unflatten(2, tensor.shape[2:])
|
||||
return tensor
|
||||
|
||||
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)
|
||||
local_mean_squared = local_mean**2
|
||||
|
||||
local_variance = F.avg_pool2d(image_tensor**2, neighborhood_size, stride=1, padding=neighborhood_size // 2, count_include_pad=False) - local_mean_squared
|
||||
|
||||
# Add a small value to prevent division by zero
|
||||
local_variance = local_variance + threshold
|
||||
|
||||
# Apply divisive normalization
|
||||
normalized_tensor = image_tensor / torch.sqrt(local_variance)
|
||||
|
||||
return normalized_tensor
|
||||
|
||||
def decorrelate_data(data):
|
||||
"""flattened = tensor.flatten(2).squeeze(0) # this code aint shit, yo
|
||||
cov_matrix = torch.cov(flattened)
|
||||
sqrt_inv_cov_matrix = torch.linalg.inv(torch.sqrt(cov_matrix))
|
||||
decorrelated_tensor = torch.dot(flattened, sqrt_inv_cov_matrix.T)
|
||||
decorrelated_tensor = decorrelated_tensor.unflatten(2, tensor.shape[2:]).unsqueeze(0)"""
|
||||
|
||||
# Reshape the 4D tensor to a 2D tensor for covariance calculation
|
||||
num_samples, num_channels, height, width = data.size()
|
||||
data_reshaped = data.view(num_samples, num_channels, -1)
|
||||
data_reshaped = data_reshaped - torch.mean(data_reshaped, dim=2, keepdim=True)
|
||||
|
||||
# Compute covariance matrix
|
||||
cov_matrix = torch.matmul(data_reshaped, data_reshaped.transpose(1, 2)) / (height * width - 1)
|
||||
|
||||
# Compute the inverse square root of the covariance matrix
|
||||
u, s, v = torch.svd(cov_matrix)
|
||||
sqrt_inv_cov_matrix = torch.matmul(u, torch.matmul(torch.diag_embed(1.0 / torch.sqrt(s)), v.transpose(1, 2)))
|
||||
|
||||
# Reshape sqrt_inv_cov_matrix to match the dimensions of data_reshaped
|
||||
sqrt_inv_cov_matrix = sqrt_inv_cov_matrix.unsqueeze(0).expand(num_samples, -1, -1, -1)
|
||||
|
||||
# Decorrelate the data
|
||||
decorrelated_data = torch.matmul(data_reshaped.transpose(1, 2), sqrt_inv_cov_matrix.transpose(2, 3))
|
||||
decorrelated_data = decorrelated_data.transpose(2, 3)
|
||||
|
||||
# Reshape back to the original shape
|
||||
decorrelated_data = decorrelated_data.view(num_samples, num_channels, height, width)
|
||||
|
||||
return decorrelated_data.to(data.device)
|
||||
|
||||
def get_low_frequency_noise(image: Tensor, threshold: float):
|
||||
# Convert image to Fourier domain
|
||||
fourier = torch.fft.fft2(image, dim=(-2, -1)) # Apply FFT along Height and Width dimensions
|
||||
|
||||
# Compute the power spectrum
|
||||
power_spectrum = torch.abs(fourier) ** 2
|
||||
|
||||
threshold = threshold ** 2
|
||||
|
||||
# Drop low-frequency components
|
||||
mask = (power_spectrum < threshold).float()
|
||||
filtered_fourier = fourier * mask
|
||||
|
||||
# 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):
|
||||
@@ -457,20 +582,25 @@ 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", "arctan", "quantile"], ),
|
||||
"tonemap_percentile": ("FLOAT", {"default": 100.0, "min": 0.0, "max": 100.0, "step": 0.05}),
|
||||
"contrast_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"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_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", "perlin", "pink", "green"], ),
|
||||
"extra_noise_type": (["gaussian", "uniform", "perlin", "pink", "green"], ),
|
||||
"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}),
|
||||
"affect_uncond": (["None", "Sharpness"], ),
|
||||
}}
|
||||
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, rescale_cfg_phi, extra_noise_type, extra_noise_method, extra_noise_multiplier):
|
||||
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):
|
||||
match sharpness_method:
|
||||
case "anisotropic":
|
||||
degrade_func = bilateral_blur
|
||||
@@ -491,6 +621,8 @@ class ModelSamplerLatentMegaModifier:
|
||||
match extra_noise_type:
|
||||
case "gaussian":
|
||||
extra_noise = torch.randn_like(cond)
|
||||
case "uniform":
|
||||
extra_noise = (torch.rand_like(cond) - 0.5) * 2 * 1.73
|
||||
case "perlin":
|
||||
cond_size_0 = cond.size(dim=2)
|
||||
cond_size_1 = cond.size(dim=3)
|
||||
@@ -513,22 +645,31 @@ class ModelSamplerLatentMegaModifier:
|
||||
std = torch.std(extra_noise)
|
||||
|
||||
extra_noise.sub_(mean).div_(std)
|
||||
|
||||
if extra_noise_lowpass > 0:
|
||||
extra_noise = get_low_frequency_noise(extra_noise, extra_noise_lowpass)
|
||||
|
||||
alpha_noise = 1.0 - (timestep / 999.0)[:, None, None, None].clone() # Get alpha multiplier, lower alpha at high sigmas/high noise
|
||||
alpha_noise *= 0.001 * extra_noise_multiplier # User-input and weaken the strength so we don't annihilate the latent.
|
||||
match extra_noise_method:
|
||||
case "add":
|
||||
cond = cond + extra_noise * alpha_noise
|
||||
uncond = uncond - extra_noise * alpha_noise
|
||||
case "add_scaled":
|
||||
cond = cond + train_difference(cond, extra_noise, cond) * alpha_noise
|
||||
uncond = uncond - train_difference(uncond, extra_noise, uncond) * alpha_noise
|
||||
case "speckle":
|
||||
cond = cond + cond * extra_noise * alpha_noise
|
||||
uncond = uncond - uncond * extra_noise * alpha_noise
|
||||
case _:
|
||||
cond = cond + extra_noise * alpha_noise
|
||||
print("Haven't heard of a noise method named like that before... (Couldn't find method)")
|
||||
|
||||
# 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
|
||||
if affect_uncond == "Sharpness":
|
||||
uncond = uncond + (uncond - degrade_func(uncond)) * alpha
|
||||
noise_pred_degraded = (degraded_cond - uncond) # New noise pred
|
||||
|
||||
# After this point, we use `noise_pred_degraded` instead of just `cond` for the final set of calculations
|
||||
@@ -552,6 +693,22 @@ class ModelSamplerLatentMegaModifier:
|
||||
new_magnitude *= top
|
||||
|
||||
noise_pred_degraded *= new_magnitude
|
||||
case "reinhard_perchannel": # Testing the flatten strategy
|
||||
flattened = noise_pred_degraded.flatten(2)
|
||||
noise_pred_vector_magnitude = (torch.linalg.vector_norm(flattened, dim=(2), keepdim=True) + 0.0000000001)
|
||||
flattened /= noise_pred_vector_magnitude
|
||||
|
||||
mean = torch.mean(noise_pred_vector_magnitude, dim=(2), keepdim=True)
|
||||
|
||||
top = (3 * (100 / tonemap_percentile) + mean) * tonemap_multiplier
|
||||
|
||||
noise_pred_vector_magnitude *= (1.0 / top)
|
||||
|
||||
new_magnitude = noise_pred_vector_magnitude / (noise_pred_vector_magnitude + 1.0)
|
||||
new_magnitude *= top
|
||||
|
||||
flattened *= new_magnitude
|
||||
noise_pred_degraded = flattened.unflatten(2, noise_pred_degraded.shape[2:])
|
||||
case "arctan":
|
||||
noise_pred_vector_magnitude = (torch.linalg.vector_norm(noise_pred_degraded, dim=(1)) + 0.0000000001)[:,None]
|
||||
noise_pred_degraded /= noise_pred_vector_magnitude
|
||||
@@ -568,13 +725,59 @@ class ModelSamplerLatentMegaModifier:
|
||||
s.clamp_(min = 1.)
|
||||
s = s.reshape(*s.shape, 1, 1, 1)
|
||||
noise_pred_degraded = noise_pred_degraded.clamp(-s, s) / s
|
||||
case "gated": # https://birchlabs.co.uk/machine-learning#dynamic-thresholding-latents so based,.,.,....,
|
||||
latent_scale = model.model.latent_format.scale_factor
|
||||
|
||||
latent = uncond + noise_pred_degraded * cond_scale # Get full latent from CFG formula
|
||||
latent /= latent_scale # Divide full CFG by latent scale (~0.13 for sdxl)
|
||||
flattened = latent.flatten(2)
|
||||
means = flattened.mean(dim=2).unsqueeze(2)
|
||||
centered_magnitudes = (flattened - means).abs().max() # Get highest magnitude of full CFG
|
||||
|
||||
flattened_pred = (noise_pred_degraded / latent_scale).flatten(2)
|
||||
|
||||
floor = 3.0560
|
||||
ceil = 42. * tonemap_multiplier # as is the answer to life, unless you modify the multiplier cuz u aint a believer in life
|
||||
|
||||
|
||||
thresholded_latent = dyn_thresh_gate(flattened_pred, centered_magnitudes, tonemap_percentile / 100., floor, ceil) # Threshold if passes ceil
|
||||
thresholded_latent = thresholded_latent.unflatten(2, noise_pred_degraded.shape[2:])
|
||||
noise_pred_degraded = thresholded_latent * latent_scale # Rescale by latent
|
||||
case "cfg-mimic":
|
||||
latent = noise_pred_degraded
|
||||
|
||||
mimic_latent = noise_pred_degraded * tonemap_multiplier
|
||||
mimic_flattened = mimic_latent.flatten(2)
|
||||
mimic_means = mimic_flattened.mean(dim=2).unsqueeze(2)
|
||||
mimic_recentered = mimic_flattened - mimic_means
|
||||
mimic_abs = mimic_recentered.abs()
|
||||
mimic_max = mimic_abs.max(dim=2).values.unsqueeze(2)
|
||||
|
||||
latent_flattened = latent.flatten(2)
|
||||
latent_means = latent_flattened.mean(dim=2).unsqueeze(2)
|
||||
latent_recentered = latent_flattened - latent_means
|
||||
latent_abs = latent_recentered.abs()
|
||||
latent_q = torch.quantile(latent_abs, tonemap_percentile / 100., dim=2).unsqueeze(2)
|
||||
s = torch.maximum(latent_q, mimic_max)
|
||||
pred_clamped = noise_pred_degraded.flatten(2).clamp(-s, s)
|
||||
pred_normalized = pred_clamped / s
|
||||
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 _:
|
||||
print("Could not tonemap, for the method was not found.")
|
||||
|
||||
# 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() # Get alpha multiplier, lower alpha at high sigmas/high noise
|
||||
alpha *= 0.001 * contrast_multiplier # User-input and weaken the strength so we don't annihilate the latent.
|
||||
noise_pred_degraded = contrast(noise_pred_degraded) * alpha + noise_pred_degraded * (1.0 - alpha) # Mix the modified latent with the existing latent by the alpha
|
||||
if contrast_multiplier > 0:
|
||||
contrast_func = contrast
|
||||
# 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)
|
||||
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)
|
||||
|
||||
# Rescale CFG
|
||||
if rescale_cfg_phi == 0:
|
||||
@@ -587,6 +790,22 @@ 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:
|
||||
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)
|
||||
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:
|
||||
match combat_method:
|
||||
case "subtract":
|
||||
combat_drift_func = center_latent_perchannel
|
||||
alpha = combat_cfg_drift
|
||||
case "subtract_w_magnitudes":
|
||||
combat_drift_func = center_latent_perchannel_with_decorrelate
|
||||
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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user