Add files via upload

This commit is contained in:
Extraltodeus
2024-04-17 16:26:09 +02:00
committed by GitHub
parent 374d62911a
commit 79feab6424
3 changed files with 323 additions and 113 deletions
+4 -1
View File
@@ -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,
}
+144 -112
View File
@@ -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, )
+175
View File
@@ -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, )