Add files via upload
This commit is contained in:
+4
-1
@@ -1,8 +1,11 @@
|
||||
from .nodes import *
|
||||
from .nodes_sag_custom import *
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Automatic CFG": simpleDynamicCFG,
|
||||
"Automatic CFG - Negative": simpleDynamicCFGlerpUncond,
|
||||
"Automatic CFG - No uncond": simpleDynamicCFGNoUncond,
|
||||
"Automatic CFG - Advanced settings": advancedDynamicCFG,
|
||||
"Automatic CFG - Advanced": advancedDynamicCFG,
|
||||
"Automatic CFG - Post rescale only": postCFGrescaleOnly,
|
||||
"SAG delayed activation": SelfAttentionGuidanceCustom,
|
||||
}
|
||||
|
||||
@@ -4,15 +4,13 @@ import torch
|
||||
import math
|
||||
|
||||
original_sampling_function = deepcopy(comfy.samplers.sampling_function)
|
||||
minimum_sigma_to_disable_uncond = 1
|
||||
minimum_sigma_to_disable_uncond = 0
|
||||
maximum_sigma_to_enable_uncond = 1000000
|
||||
no_uncond_at_all = False
|
||||
global_skip_uncond = False
|
||||
|
||||
def sampling_function_patched(model, x, timestep, uncond, cond, cond_scale, model_options={}, seed=None):
|
||||
if math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False or timestep[0] <= minimum_sigma_to_disable_uncond or no_uncond_at_all or timestep[0] > maximum_sigma_to_enable_uncond:
|
||||
if math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False or ((timestep[0] < minimum_sigma_to_disable_uncond or timestep[0] > maximum_sigma_to_enable_uncond) and global_skip_uncond):
|
||||
uncond_ = None
|
||||
if not no_uncond_at_all:
|
||||
cond_scale = 1
|
||||
else:
|
||||
uncond_ = uncond
|
||||
|
||||
@@ -48,6 +46,32 @@ def center_latent_mean_values(latent, per_channel, mult):
|
||||
latent[b] -= latent[b].mean() * mult
|
||||
return latent
|
||||
|
||||
def get_denoised_ranges(latent, measure="hard", top_k=0.25):
|
||||
chans = []
|
||||
for x in range(len(latent)):
|
||||
|
||||
max_values = torch.topk(latent[x] - latent[x].mean() if measure == "range" else latent[x], k=int(len(latent[x])*top_k), largest=True).values
|
||||
min_values = torch.topk(latent[x] - latent[x].mean() if measure == "range" else latent[x], k=int(len(latent[x])*top_k), largest=False).values
|
||||
max_val = torch.mean(max_values).item()
|
||||
min_val = torch.mean(torch.abs(min_values)).item() if (measure == "hard" or measure == "range") else abs(torch.mean(min_values).item())
|
||||
denoised_range = (max_val + min_val) / 2
|
||||
chans.append(denoised_range)
|
||||
return chans
|
||||
|
||||
def get_sigmin_sigmax(model):
|
||||
model_sampling = model.model.model_sampling
|
||||
sigmin = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_min))
|
||||
sigmax = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max))
|
||||
return sigmin, sigmax
|
||||
|
||||
def get_sigmas_start_end(sigmin, sigmax, start_percentage, end_percentage):
|
||||
high_sigma_threshold = (sigmax - sigmin) / 100 * start_percentage
|
||||
low_sigma_threshold = (sigmax - sigmin) / 100 * end_percentage
|
||||
return high_sigma_threshold, low_sigma_threshold
|
||||
|
||||
def check_skip(sigma, high_sigma_threshold, low_sigma_threshold):
|
||||
return sigma > high_sigma_threshold or sigma < low_sigma_threshold
|
||||
|
||||
class advancedDynamicCFG:
|
||||
def __init__(self):
|
||||
self.last_cfg_ht_one = 8
|
||||
@@ -56,150 +80,119 @@ class advancedDynamicCFG:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ("MODEL",),
|
||||
"center_mean_post_cfg" : ("BOOLEAN", {"default": True}),
|
||||
"center_mean_to_sigma" : ("BOOLEAN", {"default": False}),
|
||||
"automatic_cfg" : (["None","soft","hard","progressive","include_boost"], {"default": "hard"},),
|
||||
"sigma_boost" : ("BOOLEAN", {"default": True}),
|
||||
"sigma_boost_percentage": ("FLOAT", {"default": 6.86, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
|
||||
"automatic_cfg" : (["None","soft","hard","range"], {"default": "hard"},),
|
||||
|
||||
"skip_uncond" : ("BOOLEAN", {"default": True}),
|
||||
"uncond_sigma_start": ("FLOAT", {"default": 50.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"uncond_sigma_end": ("FLOAT", {"default": 6.86, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
|
||||
"lerp_uncond" : ("BOOLEAN", {"default": False}),
|
||||
"lerp_uncond_strength": ("FLOAT", {"default": 1, "min": 0.0, "max": 10.0, "step": 0.1, "round": 0.1}),
|
||||
"post_cfg_scale" : ("BOOLEAN", {"default": False}),
|
||||
"post_cfg_scale_value": ("FLOAT", {"default": 0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.1}),
|
||||
"no_uncond_mode" : ("BOOLEAN", {"default": False}),
|
||||
"uncond_start_percentage": ("FLOAT", {"default": 100.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"debug_print" : ("BOOLEAN", {"default": False}),
|
||||
"lerp_uncond_strength": ("FLOAT", {"default": 1, "min": 0.0, "max": 10.0, "step": 0.1, "round": 0.1}),
|
||||
"lerp_uncond_sigma_start": ("FLOAT", {"default": 100, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"lerp_uncond_sigma_end": ("FLOAT", {"default": 6.86, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
|
||||
"subtract_latent_mean" : ("BOOLEAN", {"default": False}),
|
||||
"subtract_latent_mean_sigma_start": ("FLOAT", {"default": 100, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"subtract_latent_mean_sigma_end": ("FLOAT", {"default": 99.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
|
||||
"latent_intensity_rescale" : ("BOOLEAN", {"default": True}),
|
||||
"latent_intensity_rescale_method" : (["soft","hard","range"], {"default": "hard"},),
|
||||
"latent_intensity_rescale_cfg" : ("FLOAT", {"default": 7.6, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
|
||||
"latent_intensity_rescale_sigma_start": ("FLOAT", {"default": 100, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"latent_intensity_rescale_sigma_end": ("FLOAT", {"default": 50, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
def patch(self, model, center_mean_post_cfg, center_mean_to_sigma,
|
||||
automatic_cfg, sigma_boost, sigma_boost_percentage, lerp_uncond=False, lerp_uncond_strength=1,
|
||||
post_cfg_scale=False, post_cfg_scale_value=8, no_uncond_mode=False, uncond_start_percentage=100, debug_print=False):
|
||||
def patch(self, model, automatic_cfg = "None",
|
||||
skip_uncond = False, uncond_sigma_start = 50, uncond_sigma_end = 6.86,
|
||||
lerp_uncond = False, lerp_uncond_strength = 1, lerp_uncond_sigma_start = 100, lerp_uncond_sigma_end = 6.86,
|
||||
subtract_latent_mean = False, subtract_latent_mean_sigma_start = 100, subtract_latent_mean_sigma_end = 99,
|
||||
latent_intensity_rescale = False, latent_intensity_rescale_sigma_start = 100, latent_intensity_rescale_sigma_end = 50,
|
||||
latent_intensity_rescale_cfg = 8, latent_intensity_rescale_method = "hard",
|
||||
ignore_pre_cfg_func = False):
|
||||
|
||||
global minimum_sigma_to_disable_uncond, maximum_sigma_to_enable_uncond, no_uncond_at_all
|
||||
no_uncond_at_all = no_uncond_mode
|
||||
model_sampling = model.model.model_sampling
|
||||
sigmin = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_min))
|
||||
sigmax = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max))
|
||||
high_sigma_threshold = (sigmax - sigmin) / 100 * uncond_start_percentage
|
||||
low_sigma_threshold = (sigmax - sigmin) / 100 * sigma_boost_percentage
|
||||
if sigma_boost_percentage > 0 and sigma_boost:
|
||||
minimum_sigma_to_disable_uncond = low_sigma_threshold
|
||||
maximum_sigma_to_enable_uncond = high_sigma_threshold
|
||||
global minimum_sigma_to_disable_uncond, maximum_sigma_to_enable_uncond, global_skip_uncond
|
||||
sigmin, sigmax = get_sigmin_sigmax(model)
|
||||
maximum_sigma_to_enable_uncond, minimum_sigma_to_disable_uncond = get_sigmas_start_end(sigmin, sigmax, uncond_sigma_start, uncond_sigma_end)
|
||||
lerp_start, lerp_end = get_sigmas_start_end(sigmin, sigmax, lerp_uncond_sigma_start, lerp_uncond_sigma_end)
|
||||
subtract_start, subtract_end = get_sigmas_start_end(sigmin, sigmax, subtract_latent_mean_sigma_start, subtract_latent_mean_sigma_end)
|
||||
rescale_start, rescale_end = get_sigmas_start_end(sigmin, sigmax, latent_intensity_rescale_sigma_start, latent_intensity_rescale_sigma_end)
|
||||
|
||||
if skip_uncond:
|
||||
global_skip_uncond = skip_uncond
|
||||
comfy.samplers.sampling_function = sampling_function_patched
|
||||
print(f"Sampling function patched. Trigger when sigmas are at: {round(minimum_sigma_to_disable_uncond.item(),4)}")
|
||||
else:
|
||||
print(f"Sampling function patched. Uncond enabled from {round(maximum_sigma_to_enable_uncond.item(),2)} to {round(minimum_sigma_to_disable_uncond.item(),2)}")
|
||||
elif not ignore_pre_cfg_func:
|
||||
global_skip_uncond = skip_uncond # just in case of mixup with another node
|
||||
comfy.samplers.sampling_function = original_sampling_function
|
||||
print(f"Sampling function unpatched.")
|
||||
|
||||
top_k = 0.25
|
||||
reference_cfg = 8
|
||||
def linear_cfg(args):
|
||||
def automatic_cfg(args):
|
||||
cond_scale = args["cond_scale"]
|
||||
input_x = args["input"]
|
||||
cond_pred = args["cond_denoised"]
|
||||
uncond_pred = args["uncond_denoised"]
|
||||
sigma = args["sigma"][0]
|
||||
|
||||
if lerp_uncond:
|
||||
lerp_weight = lerp_uncond_strength if lerp_uncond_strength > 0 else max(sigma.item(), 1)
|
||||
if lerp_weight != 1:
|
||||
uncond_pred = torch.lerp(cond_pred, uncond_pred, lerp_weight)
|
||||
cond = input_x - cond_pred
|
||||
uncond = input_x - uncond_pred
|
||||
|
||||
if no_uncond_mode:
|
||||
self.last_cfg_ht_one = cond_scale
|
||||
return cond
|
||||
|
||||
if sigma == sigmax or cond_scale > 1:
|
||||
self.last_cfg_ht_one = cond_scale
|
||||
|
||||
target_intensity = self.last_cfg_ht_one / 10
|
||||
|
||||
if sigma_boost and cond_scale > 1:
|
||||
for b in range(len(cond)):
|
||||
for c in range(len(cond[b])):
|
||||
uncond[b][c] = uncond[b][c] * torch.norm(cond[b][c]) / torch.norm(uncond[b][c])
|
||||
if (check_skip(sigma, maximum_sigma_to_enable_uncond, minimum_sigma_to_disable_uncond) and skip_uncond) or cond_scale == 1:
|
||||
return input_x - cond_pred
|
||||
|
||||
if automatic_cfg == "None" or (cond_scale == 1 and automatic_cfg != "include_boost"):
|
||||
if lerp_uncond and not check_skip(sigma, lerp_start, lerp_end) and lerp_uncond_strength != 1:
|
||||
uncond_pred = torch.lerp(cond_pred, uncond_pred, lerp_uncond_strength)
|
||||
cond = input_x - cond_pred
|
||||
uncond = input_x - uncond_pred
|
||||
|
||||
if automatic_cfg == "None":
|
||||
return uncond + cond_scale * (cond - uncond)
|
||||
|
||||
if cond_scale > 1:
|
||||
denoised_tmp = input_x - (uncond + reference_cfg * (cond - uncond))
|
||||
else:
|
||||
denoised_tmp = input_x + cond_pred
|
||||
|
||||
denoised_tmp = input_x - (uncond + reference_cfg * (cond - uncond))
|
||||
|
||||
for b in range(len(denoised_tmp)):
|
||||
denoised_ranges = get_denoised_ranges(denoised_tmp[b], automatic_cfg, top_k)
|
||||
for c in range(len(denoised_tmp[b])):
|
||||
channel = denoised_tmp[b][c]
|
||||
max_values = torch.topk(channel, k=int(len(channel)*top_k), largest=True ).values
|
||||
min_values = torch.topk(channel, k=int(len(channel)*top_k), largest=False).values
|
||||
max_val = torch.mean(max_values).item()
|
||||
|
||||
if automatic_cfg == "soft":
|
||||
min_val = abs(torch.mean(min_values).item())
|
||||
elif automatic_cfg == "hard" or automatic_cfg == "include_boost":
|
||||
min_val = torch.mean(torch.abs(min_values)).item()
|
||||
elif automatic_cfg == "progressive":
|
||||
min_val = torch.mean(torch.abs(min_values)).item()
|
||||
s_progression = map_sigma(sigma, sigmax, sigmin)
|
||||
target_intensity = 1.1 * target_intensity * s_progression + 0.9 * target_intensity * (1 - s_progression)
|
||||
|
||||
denoised_range = (max_val + min_val) / 2
|
||||
scale_correction = target_intensity / denoised_range
|
||||
tmp_scale = reference_cfg * scale_correction
|
||||
|
||||
if debug_print:
|
||||
print(f"c{c}: {tmp_scale} / {scale_correction}")
|
||||
print(f"denoised_range: {denoised_range}")
|
||||
|
||||
if cond_scale > 1:
|
||||
denoised_tmp[b][c] = uncond[b][c] + tmp_scale * (cond[b][c] - uncond[b][c])
|
||||
else:
|
||||
denoised_tmp[b][c] = scale_correction * cond[b][c]
|
||||
|
||||
# The scaling has been done per channel, now we set it back to norm.
|
||||
if cond_scale == 1:
|
||||
denoised_tmp = denoised_tmp * cond.norm() / denoised_tmp.norm()
|
||||
fixeds_scale = reference_cfg * target_intensity / denoised_ranges[c]
|
||||
denoised_tmp[b][c] = uncond[b][c] + fixeds_scale * (cond[b][c] - uncond[b][c])
|
||||
|
||||
return denoised_tmp
|
||||
|
||||
def center_mean_latent_post_cfg(args):
|
||||
denoised = args["denoised"]
|
||||
sigma = args["sigma"][0]
|
||||
mult = map_sigma(sigma, sigmax, sigmin) if center_mean_to_sigma else 1
|
||||
denoised = center_latent_mean_values(denoised, False, mult)
|
||||
if check_skip(sigma, subtract_start, subtract_end):
|
||||
return denoised
|
||||
denoised = center_latent_mean_values(denoised, False, 1)
|
||||
return denoised
|
||||
|
||||
def rescale_post_cfg(args):
|
||||
denoised = args["denoised"]
|
||||
sigma = args["sigma"][0]
|
||||
if sigma <= minimum_sigma_to_disable_uncond:
|
||||
denoised = args["denoised"]
|
||||
sigma = args["sigma"][0]
|
||||
|
||||
if check_skip(sigma, rescale_start, rescale_end):
|
||||
return denoised
|
||||
target_intensity = latent_intensity_rescale_cfg / 10
|
||||
for b in range(len(denoised)):
|
||||
for c in range(len(denoised[b])): #TODO make a function for the scaling
|
||||
channel = denoised[b][c]
|
||||
max_values = torch.topk(channel, k=int(len(channel)*top_k), largest=True ).values
|
||||
min_values = torch.topk(channel, k=int(len(channel)*top_k), largest=False).values
|
||||
max_val = torch.mean(max_values).item()
|
||||
min_val = torch.mean(torch.abs(min_values)).item()
|
||||
denoised_range = (max_val + min_val) / 2
|
||||
if no_uncond_mode or post_cfg_scale_value == 0:
|
||||
target_intensity = self.last_cfg_ht_one / 10
|
||||
else:
|
||||
target_intensity = post_cfg_scale_value / 10
|
||||
scale_correction = target_intensity / denoised_range
|
||||
denoised[b][c] = channel * scale_correction
|
||||
denoised_ranges = get_denoised_ranges(denoised[b], latent_intensity_rescale_method)
|
||||
for c in range(len(denoised[b])):
|
||||
scale_correction = target_intensity / denoised_ranges[c]
|
||||
denoised[b][c] = denoised[b][c] * scale_correction
|
||||
return denoised
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_cfg_function(linear_cfg, disable_cfg1_optimization=False)
|
||||
if center_mean_post_cfg or no_uncond_mode:
|
||||
if not ignore_pre_cfg_func:
|
||||
m.set_model_sampler_cfg_function(automatic_cfg, disable_cfg1_optimization = False)
|
||||
if subtract_latent_mean:
|
||||
m.set_model_sampler_post_cfg_function(center_mean_latent_post_cfg)
|
||||
if post_cfg_scale or no_uncond_mode:
|
||||
if latent_intensity_rescale:
|
||||
m.set_model_sampler_post_cfg_function(rescale_post_cfg)
|
||||
return (m, )
|
||||
|
||||
@@ -215,9 +208,13 @@ class simpleDynamicCFG:
|
||||
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
def patch(self, model, boost, color_balance=False):
|
||||
def patch(self, model, boost):
|
||||
advcfg = advancedDynamicCFG()
|
||||
m = advcfg.patch(model,color_balance,color_balance,"hard" if boost else "soft", boost, 6.86)[0]
|
||||
m = advcfg.patch(model,
|
||||
skip_uncond = boost,
|
||||
uncond_sigma_start = 100, uncond_sigma_end = 6.86,
|
||||
automatic_cfg = "hard" if boost else "soft"
|
||||
)[0]
|
||||
return (m, )
|
||||
|
||||
class simpleDynamicCFGlerpUncond:
|
||||
@@ -235,10 +232,12 @@ class simpleDynamicCFGlerpUncond:
|
||||
|
||||
def patch(self, model, boost, negative_strength):
|
||||
advcfg = advancedDynamicCFG()
|
||||
# automatic_cfg="progressive" if negative_strength == 1 else "hard"
|
||||
m = advcfg.patch(model=model, center_mean_post_cfg=False, center_mean_to_sigma=False,
|
||||
automatic_cfg="hard", sigma_boost=boost, sigma_boost_percentage=6.86,
|
||||
lerp_uncond=negative_strength != 1, lerp_uncond_strength=negative_strength)[0]
|
||||
m = advcfg.patch(model=model,
|
||||
automatic_cfg="hard", sigma_boost=boost,
|
||||
uncond_sigma_start = 100, uncond_sigma_end = 6.86,
|
||||
lerp_uncond=negative_strength != 1, lerp_uncond_strength=negative_strength,
|
||||
lerp_uncond_sigma_start = 100, lerp_uncond_sigma_end = 6.86
|
||||
)[0]
|
||||
return (m, )
|
||||
|
||||
class simpleDynamicCFGNoUncond:
|
||||
@@ -255,6 +254,39 @@ class simpleDynamicCFGNoUncond:
|
||||
def patch(self, model):
|
||||
advcfg = advancedDynamicCFG()
|
||||
m = advcfg.patch(model=model, center_mean_post_cfg=True, center_mean_to_sigma=False,
|
||||
automatic_cfg="None", sigma_boost="None", sigma_boost_percentage=6.86,
|
||||
automatic_cfg="None", sigma_boost=False, sigma_boost_percentage=6.86,
|
||||
no_uncond_mode=True)[0]
|
||||
return (m, )
|
||||
|
||||
class postCFGrescaleOnly:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ("MODEL",),
|
||||
"subtract_latent_mean" : ("BOOLEAN", {"default": True}),
|
||||
"subtract_latent_mean_sigma_start": ("FLOAT", {"default": 100, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"subtract_latent_mean_sigma_end": ("FLOAT", {"default": 99.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"latent_intensity_rescale" : ("BOOLEAN", {"default": True}),
|
||||
"latent_intensity_rescale_method" : (["soft","hard","range"], {"default": "hard"},),
|
||||
"latent_intensity_rescale_cfg" : ("FLOAT", {"default": 8, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
|
||||
"latent_intensity_rescale_sigma_start": ("FLOAT", {"default": 100, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"latent_intensity_rescale_sigma_end": ("FLOAT", {"default": 50, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
def patch(self, model,
|
||||
subtract_latent_mean, subtract_latent_mean_sigma_start, subtract_latent_mean_sigma_end,
|
||||
latent_intensity_rescale, latent_intensity_rescale_method, latent_intensity_rescale_cfg, latent_intensity_rescale_sigma_start, latent_intensity_rescale_sigma_end
|
||||
):
|
||||
advcfg = advancedDynamicCFG()
|
||||
m = advcfg.patch(model=model,
|
||||
subtract_latent_mean = subtract_latent_mean,
|
||||
subtract_latent_mean_sigma_start = subtract_latent_mean_sigma_start, subtract_latent_mean_sigma_end = subtract_latent_mean_sigma_end,
|
||||
latent_intensity_rescale = latent_intensity_rescale, latent_intensity_rescale_cfg = latent_intensity_rescale_cfg, latent_intensity_rescale_method = latent_intensity_rescale_method,
|
||||
latent_intensity_rescale_sigma_start = latent_intensity_rescale_sigma_start, latent_intensity_rescale_sigma_end = latent_intensity_rescale_sigma_end,
|
||||
ignore_pre_cfg_func = True
|
||||
)[0]
|
||||
return (m, )
|
||||
@@ -0,0 +1,175 @@
|
||||
import torch
|
||||
from torch import einsum
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
|
||||
from einops import rearrange, repeat
|
||||
import os
|
||||
from comfy.ldm.modules.attention import optimized_attention, _ATTN_PRECISION
|
||||
import comfy.samplers
|
||||
|
||||
# from comfy/ldm/modules/attention.py
|
||||
# but modified to return attention scores as well as output
|
||||
def attention_basic_with_sim(q, k, v, heads, mask=None):
|
||||
b, _, dim_head = q.shape
|
||||
dim_head //= heads
|
||||
scale = dim_head ** -0.5
|
||||
|
||||
h = heads
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(b, -1, heads, dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b * heads, -1, dim_head)
|
||||
.contiguous(),
|
||||
(q, k, v),
|
||||
)
|
||||
|
||||
# force cast to fp32 to avoid overflowing
|
||||
if _ATTN_PRECISION =="fp32":
|
||||
sim = einsum('b i d, b j d -> b i j', q.float(), k.float()) * scale
|
||||
else:
|
||||
sim = einsum('b i d, b j d -> b i j', q, k) * scale
|
||||
|
||||
del q, k
|
||||
|
||||
if mask is not None:
|
||||
mask = rearrange(mask, 'b ... -> b (...)')
|
||||
max_neg_value = -torch.finfo(sim.dtype).max
|
||||
mask = repeat(mask, 'b j -> (b h) () j', h=h)
|
||||
sim.masked_fill_(~mask, max_neg_value)
|
||||
|
||||
# attention, what we cannot get enough of
|
||||
sim = sim.softmax(dim=-1)
|
||||
|
||||
out = einsum('b i j, b j d -> b i d', sim.to(v.dtype), v)
|
||||
out = (
|
||||
out.unsqueeze(0)
|
||||
.reshape(b, heads, -1, dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b, -1, heads * dim_head)
|
||||
)
|
||||
return (out, sim)
|
||||
|
||||
def create_blur_map(x0, attn, sigma=3.0, threshold=1.0):
|
||||
# reshape and GAP the attention map
|
||||
_, hw1, hw2 = attn.shape
|
||||
b, _, lh, lw = x0.shape
|
||||
attn = attn.reshape(b, -1, hw1, hw2)
|
||||
# Global Average Pool
|
||||
mask = attn.mean(1, keepdim=False).sum(1, keepdim=False) > threshold
|
||||
ratio = 2**(math.ceil(math.sqrt(lh * lw / hw1)) - 1).bit_length()
|
||||
mid_shape = [math.ceil(lh / ratio), math.ceil(lw / ratio)]
|
||||
|
||||
# Reshape
|
||||
mask = (
|
||||
mask.reshape(b, *mid_shape)
|
||||
.unsqueeze(1)
|
||||
.type(attn.dtype)
|
||||
)
|
||||
# Upsample
|
||||
mask = F.interpolate(mask, (lh, lw))
|
||||
|
||||
blurred = gaussian_blur_2d(x0, kernel_size=9, sigma=sigma)
|
||||
blurred = blurred * mask + x0 * (1 - mask)
|
||||
return blurred
|
||||
|
||||
def gaussian_blur_2d(img, kernel_size, sigma):
|
||||
ksize_half = (kernel_size - 1) * 0.5
|
||||
|
||||
x = torch.linspace(-ksize_half, ksize_half, steps=kernel_size)
|
||||
|
||||
pdf = torch.exp(-0.5 * (x / sigma).pow(2))
|
||||
|
||||
x_kernel = pdf / pdf.sum()
|
||||
x_kernel = x_kernel.to(device=img.device, dtype=img.dtype)
|
||||
|
||||
kernel2d = torch.mm(x_kernel[:, None], x_kernel[None, :])
|
||||
kernel2d = kernel2d.expand(img.shape[-3], 1, kernel2d.shape[0], kernel2d.shape[1])
|
||||
|
||||
padding = [kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2]
|
||||
|
||||
img = F.pad(img, padding, mode="reflect")
|
||||
img = F.conv2d(img, kernel2d, groups=img.shape[-3])
|
||||
return img
|
||||
|
||||
class SelfAttentionGuidanceCustom:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"scale": ("FLOAT", {"default": 0.5, "min": -2.0, "max": 5.0, "step": 0.1}),
|
||||
"blur_sigma": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
||||
"sigma_start_percentage": ("FLOAT", {"default": 100.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
"sigma_end_percentage": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
def patch(self, model, scale, blur_sigma, sigma_start_percentage, sigma_end_percentage):
|
||||
m = model.clone()
|
||||
|
||||
model_sampling = model.model.model_sampling
|
||||
sigmin = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_min))
|
||||
sigmax = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max))
|
||||
high_sigma_threshold = (sigmax - sigmin) / 100 * sigma_start_percentage
|
||||
low_sigma_threshold = (sigmax - sigmin) / 100 * sigma_end_percentage
|
||||
|
||||
attn_scores = None
|
||||
|
||||
# TODO: make this work properly with chunked batches
|
||||
# currently, we can only save the attn from one UNet call
|
||||
def attn_and_record(q, k, v, extra_options):
|
||||
nonlocal attn_scores
|
||||
# if uncond, save the attention scores
|
||||
heads = extra_options["n_heads"]
|
||||
cond_or_uncond = extra_options["cond_or_uncond"]
|
||||
b = q.shape[0] // len(cond_or_uncond)
|
||||
if 1 in cond_or_uncond:
|
||||
uncond_index = cond_or_uncond.index(1)
|
||||
# do the entire attention operation, but save the attention scores to attn_scores
|
||||
(out, sim) = attention_basic_with_sim(q, k, v, heads=heads)
|
||||
# when using a higher batch size, I BELIEVE the result batch dimension is [uc1, ... ucn, c1, ... cn]
|
||||
n_slices = heads * b
|
||||
attn_scores = sim[n_slices * uncond_index:n_slices * (uncond_index+1)]
|
||||
return out
|
||||
else:
|
||||
return optimized_attention(q, k, v, heads=heads)
|
||||
|
||||
def post_cfg_function(args):
|
||||
nonlocal attn_scores
|
||||
uncond_attn = attn_scores
|
||||
|
||||
sag_scale = scale
|
||||
sag_sigma = blur_sigma
|
||||
sag_threshold = 1.0
|
||||
model = args["model"]
|
||||
uncond_pred = args["uncond_denoised"]
|
||||
uncond = args["uncond"]
|
||||
cfg_result = args["denoised"]
|
||||
sigma = args["sigma"]
|
||||
model_options = args["model_options"]
|
||||
x = args["input"]
|
||||
if not isinstance(uncond, torch.Tensor):
|
||||
return cfg_result
|
||||
if min(cfg_result.shape[2:]) <= 4: #skip when too small to add padding
|
||||
return cfg_result
|
||||
if sigma[0] > high_sigma_threshold or sigma[0] < low_sigma_threshold:
|
||||
return cfg_result
|
||||
# create the adversarially blurred image
|
||||
degraded = create_blur_map(uncond_pred, uncond_attn, sag_sigma, sag_threshold)
|
||||
degraded_noised = degraded + x - uncond_pred
|
||||
# call into the UNet
|
||||
(sag, _) = comfy.samplers.calc_cond_batch(model, [uncond, None], degraded_noised, sigma, model_options)
|
||||
# comfy.samplers.calc_cond_uncond_batch(model, uncond, None, degraded_noised, sigma, model_options)
|
||||
|
||||
return cfg_result + (degraded - sag) * sag_scale
|
||||
|
||||
m.set_model_sampler_post_cfg_function(post_cfg_function, disable_cfg1_optimization=False)
|
||||
|
||||
# from diffusers:
|
||||
# unet.mid_block.attentions[0].transformer_blocks[0].attn1.patch
|
||||
m.set_model_attn1_replace(attn_and_record, "middle", 0, 0)
|
||||
|
||||
return (m, )
|
||||
Reference in New Issue
Block a user