Files

236 lines
9.6 KiB
Python

import torch
from copy import deepcopy
from comfy.model_management import interrupt_current_processing
from math import floor
selfnorm = lambda x: x / x.norm()
def topk_average(latent, top_k=0.25):
max_values = torch.topk(latent, k=int(len(latent)*top_k), largest=True).values
min_values = torch.topk(latent, k=int(len(latent)*top_k), largest=False).values
max_val = torch.mean(max_values).item()
min_val = torch.mean(torch.abs(min_values)).item()
value_range = (max_val + min_val) / 2
return value_range
def normalized_pow(t,p):
t_norm = t.norm()
t_sign = t.sign()
t_pow = (t / t_norm).abs().pow(p)
t_pow = selfnorm(t_pow) * t_norm * t_sign
return t_pow
def normalize_adjust(a,b,strength=1):
norm_a = torch.linalg.norm(a)
a = selfnorm(a)
b = selfnorm(b)
res = b - a * (a * b).sum()
if res.isnan().any():
res = torch.nan_to_num(res, nan=0.0)
a = a - res * strength
return a * norm_a
class uncondZeroNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"scale": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 10.0, "step": 1/20, "round": 0.01}),
"pre_fix" : ("BOOLEAN", {"default": True}),
"pre_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 1/10, "round": 0.1}),
# "exp_fix" : ("BOOLEAN", {"default": False}),
# "exp_scale": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 2.0, "step": 1/20, "round": 0.01}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches"
def patch(self, model, scale, pre_fix, pre_scale, exp_fix=False, exp_scale=1):
model_sampling = model.model.model_sampling
sigma_max = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)).item()
prev_cond = None
prev_uncond = None
def cfg_or_zero_wrapper(args):
nonlocal prev_cond, prev_uncond
if args["sigma"][0].item() > (sigma_max - 1):
prev_cond = None
prev_uncond = None
if torch.any(args['uncond_denoised']):
return automatic_cfg(args)
return uncond_zero(args)
def automatic_cfg(args):
nonlocal prev_cond, prev_uncond
cond = args["cond_denoised"]
uncond = args["uncond_denoised"]
x_orig = args["input"]
cond_scale = args["cond_scale"]
result = torch.zeros_like(x_orig)
for b in range(len(x_orig)):
for c in range(len(cond[b])):
mes = topk_average(8 * cond[b][c] - 7 * uncond[b][c])
result[b][c] = (x_orig[b][c] - uncond[b][c]) + ((x_orig[b][c] - cond[b][c]) - (x_orig[b][c] - uncond[b][c])) * 8 * (cond_scale / 10) / mes
prev_cond = cond
prev_uncond = uncond
return result
def uncond_zero(args):
nonlocal prev_cond, prev_uncond
cond = args["cond_denoised"]
x_orig = args["input"]
sigma = args["sigma"][0].item()
first_step = True
if sigma <= 1:
return x_orig - cond
if sigma < (sigma_max - 1):
first_step = False
result = torch.zeros_like(x_orig)
for b in range(len(x_orig)):
if exp_fix and exp_scale != 1:
cond[b] = normalized_pow(cond[b], exp_scale)
for c in range(len(cond[b])):
if not first_step and pre_fix:
cond[b][c] = normalize_adjust(cond[b][c], prev_cond[b][c], pre_scale)
mes = topk_average(cond[b][c]) ** 0.5 # the square root is to dampen the variations
result[b][c] = x_orig[b][c] - cond[b][c] * scale / mes
prev_cond = cond
return result
m = model.clone()
m.set_model_sampler_cfg_function(cfg_or_zero_wrapper)
return (m, )
def sub_neg_to_pos(a, b, uncond_strength):
norm_a = torch.linalg.norm(a)
res = b - a * (a / norm_a * (b / norm_a)).sum()
if res.isnan().any():
res = torch.nan_to_num(res, nan=0.0)
a = a - res * uncond_strength
return a, res
# While this may look weird, among 26 different going from thoughtful to complete nonsense, this gave sharper results.
def post_cond_out_wrapped(a, b, c, strength):
if torch.equal(a, c) or torch.equal(b, c):
return a, b
a_delta = a - c
b_delta = b - c
_, res_a = sub_neg_to_pos(a_delta, b_delta, 1)
_, res_b = sub_neg_to_pos(b_delta, a_delta, 1)
res_a, _ = sub_neg_to_pos(res_a, a, 1)
res_b, _ = sub_neg_to_pos(res_b, b, 1)
a, _ = sub_neg_to_pos(a, res_a, strength)
b, _ = sub_neg_to_pos(b, res_b, strength)
return a, b
def post_cond_out(a, b, c, strength):
for x in range(floor(strength)):
a, b = post_cond_out_wrapped(a,b,c,1)
if (strength - floor(strength)) > 0:
a, b = post_cond_out_wrapped(a,b,c,strength - floor(strength))
return a, b
class cond_combine_pos_neg:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"positive_conditioning": ("CONDITIONING", ),
"negative_conditioning": ("CONDITIONING", ),
"empty_conditioning": ("CONDITIONING",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1}),
}
}
FUNCTION = "exec"
RETURN_TYPES = ("CONDITIONING","CONDITIONING",)
RETURN_NAMES = ("positive","negative",)
CATEGORY = "conditioning"
def exec(self, positive_conditioning, negative_conditioning, empty_conditioning, strength):
if strength == 0:
return(positive_conditioning,negative_conditioning,)
cond_copy_1 = deepcopy(positive_conditioning)
cond_copy_2 = deepcopy(negative_conditioning)
cond_copy_3 = deepcopy(empty_conditioning)
s = 1
for x in range(min(len(cond_copy_1),len(cond_copy_2))):
n_cond_slices = min(cond_copy_1[x][0].shape[1],cond_copy_2[x][0].shape[1]) // 77
for n in range(n_cond_slices):
if cond_copy_1[x][0].shape[-1] == 2048:
cond_copy_1[x][0][...,n*77+s:(n+1)*77,0:768], cond_copy_2[x][0][...,n*77+s:(n+1)*77,0:768] = post_cond_out(cond_copy_1[x][0][...,n*77+s:(n+1)*77,0:768], cond_copy_2[x][0][...,n*77+s:(n+1)*77,0:768], cond_copy_3[0][0][...,s:77,0:768], strength)
cond_copy_1[x][0][...,n*77+s:(n+1)*77,768:2048], cond_copy_2[x][0][...,n*77+s:(n+1)*77,768:2048] = post_cond_out(cond_copy_1[x][0][...,n*77+s:(n+1)*77,768:2048], cond_copy_2[x][0][...,n*77+s:(n+1)*77,768:2048], cond_copy_3[0][0][...,s:77,768:2048], strength)
else:
cond_copy_1[x][0][...,n*77+s:(n+1)*77,:], cond_copy_2[x][0][...,n*77+s:(n+1)*77,:] = post_cond_out(cond_copy_1[x][0][...,n*77+s:(n+1)*77,:], cond_copy_2[x][0][...,n*77+s:(n+1)*77,:],cond_copy_3[0][0][...,s:77,:], strength)
return (cond_copy_1,cond_copy_2,)
class conditioningCropAdd:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"conditioning": ("CONDITIONING", ),
"empty_conditioning": ("CONDITIONING", ),
"context_length" : ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}),
"enabled" : ("BOOLEAN", {"default": True}),
}
}
FUNCTION = "exec"
RETURN_TYPES = ("CONDITIONING",)
RETURN_NAMES = ("CONDITIONING",)
CATEGORY = "conditioning"
def exec(self, conditioning, empty_conditioning, context_length, enabled):
if not enabled: return (conditioning, )
cond_copy = deepcopy(conditioning)
for x in range(len(cond_copy)):
n_cond_slices = cond_copy[x][0].shape[1] // 77
if n_cond_slices == context_length:
continue
elif n_cond_slices > context_length:
cropped_cond = cond_copy[x][0][...,0:context_length*77,:]
cond_copy[x][0] = cropped_cond
else:
for y in range(context_length - n_cond_slices):
cond_copy[x][0] = torch.cat((cond_copy[x][0][..., 0:(y + context_length) * 77, :], empty_conditioning[0][0][...,0:77,:]), dim=-2)
return (cond_copy,)
class interruptNaNpatch:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"replace_values" : ("BOOLEAN", {"default": True}),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches"
def patch(self, model, replace_values, **kwargs):
def interrupt_on_nan(args):
denoised = args["denoised"]
if torch.isnan(denoised).any() or torch.isinf(denoised).any():
if replace_values:
denoised = torch.nan_to_num(denoised, nan=0.0)
else:
interrupt_current_processing()
return denoised
m = model.clone()
m.set_model_sampler_post_cfg_function(interrupt_on_nan)
return (m, )