Change how input is calculated to restore previous functionality
This commit is contained in:
+127
-31
@@ -4,6 +4,59 @@ import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
|
||||
'''
|
||||
The following snippet is utilized from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/
|
||||
'''
|
||||
def min_(tensor_list):
|
||||
# return the element-wise min of the tensor list.
|
||||
x = torch.stack(tensor_list)
|
||||
mn = x.min(axis=0)[0]
|
||||
return mn#torch.clamp(mn, min=-1)
|
||||
|
||||
def max_(tensor_list):
|
||||
# return the element-wise max of the tensor list.
|
||||
x = torch.stack(tensor_list)
|
||||
mx = x.max(axis=0)[0]
|
||||
return mx#torch.clamp(mx, max=1)
|
||||
def contrast_adaptive_sharpening(image, amount):
|
||||
img = F.pad(image, pad=(1, 1, 1, 1))
|
||||
absmean = torch.abs(image.mean())
|
||||
|
||||
a = img[..., :-2, :-2]
|
||||
b = img[..., :-2, 1:-1]
|
||||
c = img[..., :-2, 2:]
|
||||
d = img[..., 1:-1, :-2]
|
||||
e = img[..., 1:-1, 1:-1]
|
||||
f = img[..., 1:-1, 2:]
|
||||
g = img[..., 2:, :-2]
|
||||
h = img[..., 2:, 1:-1]
|
||||
i = img[..., 2:, 2:]
|
||||
|
||||
# Computing contrast
|
||||
cross = (b, d, e, f, h)
|
||||
mn = min_(cross)
|
||||
mx = max_(cross)
|
||||
|
||||
diag = (a, c, g, i)
|
||||
mn2 = min_(diag)
|
||||
mx2 = max_(diag)
|
||||
mx = mx + mx2
|
||||
mn = mn + mn2
|
||||
|
||||
# Computing local weight
|
||||
inv_mx = torch.reciprocal(mx)
|
||||
amp = inv_mx * torch.minimum(mn, (2 - mx))
|
||||
|
||||
# scaling
|
||||
amp = torch.copysign(torch.sqrt(torch.abs(amp)), amp)
|
||||
w = - amp * (amount * (1/5 - 1/8) + 1/8)
|
||||
div = torch.reciprocal(1 + 4*w).clamp(-1, 1)
|
||||
|
||||
output = ((b + d + f + h)*w + e) * div
|
||||
output = torch.nan_to_num(output)
|
||||
|
||||
return (output.to(image.device))
|
||||
|
||||
'''
|
||||
The following gaussian functions were utilized from the Fooocus UI, many thanks to github.com/Illyasviel !
|
||||
'''
|
||||
@@ -485,11 +538,12 @@ def contrast(x: Tensor):
|
||||
|
||||
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)
|
||||
#mean = x.mean(dim=(2,3), keepdim=True)
|
||||
stddev = x.std(dim=(1,2,3), keepdim=True)
|
||||
diff_mean = ((x / stddev) - x).mean(dim=(1,2,3), keepdim=True)
|
||||
# Scale the pixel values by the standard deviation
|
||||
scaled_pixels = (x - mean) / stddev
|
||||
return scaled_pixels
|
||||
scaled_pixels = x / stddev
|
||||
return scaled_pixels - diff_mean
|
||||
|
||||
def center_latent(tensor): #https://birchlabs.co.uk/machine-learning#combating-mean-drift-in-cfg
|
||||
"""Centers on 0 to combat CFG drift."""
|
||||
@@ -692,19 +746,48 @@ 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
|
||||
case "dyncfg-halfcosine-mimic":
|
||||
noise_pred = conditioning - unconditioning
|
||||
|
||||
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
|
||||
|
||||
latent = noise_pred
|
||||
|
||||
mimic_latent = noise_pred * time_factor
|
||||
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, 0.995, dim=2).unsqueeze(2)
|
||||
s = torch.maximum(latent_q, mimic_max)
|
||||
pred_clamped = noise_pred.flatten(2).clamp(-s, s)
|
||||
pred_normalized = pred_clamped / s
|
||||
pred_renorm = pred_normalized * mimic_max
|
||||
pred_uncentered = pred_renorm + latent_means
|
||||
noise_pred_degraded = pred_uncentered.unflatten(2, noise_pred.shape[2:])
|
||||
|
||||
noise_pred /= noise_pred_magnitude
|
||||
|
||||
noise_pred_timescaled_magnitude = (torch.linalg.vector_norm(noise_pred_degraded, dim=(1)) + 0.0000000001)[:,None]
|
||||
noise_pred *= noise_pred_timescaled_magnitude
|
||||
return noise_pred
|
||||
|
||||
@@ -714,7 +797,7 @@ class ModelSamplerLatentMegaModifier:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"sharpness_multiplier": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"sharpness_method": (["anisotropic", "gaussian"], ),
|
||||
"sharpness_method": (["anisotropic", "joint-anisotropic", "gaussian", "cas"], ),
|
||||
"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", "spatial-norm"], ),
|
||||
"tonemap_percentile": ("FLOAT", {"default": 100.0, "min": 0.0, "max": 100.0, "step": 0.005}),
|
||||
@@ -732,7 +815,7 @@ class ModelSamplerLatentMegaModifier:
|
||||
"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"], ),
|
||||
"dyn_cfg_augmentation": (["None", "dyncfg-halfcosine", "dyncfg-halfcosine-mimic"], ),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "mega_modify"
|
||||
@@ -740,20 +823,20 @@ 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, 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
|
||||
case "gaussian":
|
||||
degrade_func = gaussian_filter_2d
|
||||
case _:
|
||||
print("For some reason, the sharpness filter could not be found.")
|
||||
|
||||
def modify_latent(args):
|
||||
x_input = args["input"]
|
||||
cond = args["cond"]
|
||||
uncond = args["uncond"]
|
||||
cond_scale = args["cond_scale"]
|
||||
timestep = model.model.model_sampling.timestep(args["timestep"])
|
||||
sigma = args["sigma"]
|
||||
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
|
||||
#print(model.model.model_sampling.timestep(timestep))
|
||||
|
||||
x = x_input / (sigma * sigma + 1.0)
|
||||
cond = ((x - (x_input - cond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma)
|
||||
uncond = ((x - (x_input - uncond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma)
|
||||
|
||||
noise_pred = (cond - uncond)
|
||||
|
||||
# Extra noise
|
||||
@@ -806,12 +889,25 @@ class ModelSamplerLatentMegaModifier:
|
||||
case _:
|
||||
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.
|
||||
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 = degrade_func(uncond) * alpha + uncond * (1.0 - alpha)
|
||||
if sharpness_multiplier > 0.0:
|
||||
match sharpness_method:
|
||||
case "anisotropic":
|
||||
degrade_func = bilateral_blur
|
||||
case "joint-anisotropic":
|
||||
s, m = torch.std_mean(args["cond"], dim=(1, 2, 3), keepdim=True)
|
||||
degrade_func = lambda img: joint_bilateral_blur(img, (args["cond"] - m) / s, 13, 3.0, 3.0, "reflect", "l1")
|
||||
case "gaussian":
|
||||
degrade_func = gaussian_filter_2d
|
||||
case "cas":
|
||||
degrade_func = lambda image: contrast_adaptive_sharpening(image, amount=sigma.clamp(max=1.00).item())
|
||||
case _:
|
||||
print("For some reason, the sharpness filter could not be found.")
|
||||
# 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.
|
||||
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 - (degrade_func(uncond) * alpha + uncond * (1.0 - alpha))
|
||||
|
||||
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
|
||||
@@ -910,8 +1006,8 @@ class ModelSamplerLatentMegaModifier:
|
||||
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)
|
||||
time = -(math.cos(time * math.pi) / (3)) + (2/3) # 0.33333 to 1.0, half cosine
|
||||
noise_pred_degraded = spatial_norm_chw_thresholding(noise_pred_degraded, (tonemap_multiplier / 2 / cond_scale) * time)
|
||||
case _:
|
||||
print("Could not tonemap, for the method was not found.")
|
||||
|
||||
@@ -932,12 +1028,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 + 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...
|
||||
noise_pred_degraded = contrast_func(noise_pred_degraded) * alpha + (noise_pred_degraded) * (1.0 - alpha) # 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 + args["input"]) * alpha + (noise_pred_degraded + args["input"]) * (1.0 - alpha)) - args["input"]
|
||||
noise_pred_degraded = contrast_func(noise_pred_degraded) * alpha + (noise_pred_degraded) * (1.0 - alpha)
|
||||
|
||||
# Rescale CFG
|
||||
if rescale_cfg_phi == 0:
|
||||
@@ -952,9 +1048,9 @@ class ModelSamplerLatentMegaModifier:
|
||||
|
||||
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 (?).
|
||||
alpha ** 0.025 # 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)
|
||||
high_noise = divisive_normalization(x_final, (divisive_norm_size * 2) + 1)
|
||||
x_final = high_noise * alpha + x_final * (1.0 - alpha)
|
||||
|
||||
if combat_cfg_drift > 0:
|
||||
@@ -973,7 +1069,7 @@ class ModelSamplerLatentMegaModifier:
|
||||
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
|
||||
return x_input - (x - x_final * sigma / (sigma * sigma + 1.0) ** 0.5) # General formula for CFG. uncond + (cond - uncond) * cond_scale
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_cfg_function(modify_latent)
|
||||
|
||||
Reference in New Issue
Block a user