Add files via upload

This commit is contained in:
Extraltodeus
2024-07-03 01:02:23 +02:00
committed by GitHub
parent 3c1c8e9daa
commit 5babf5763d
2 changed files with 188 additions and 58 deletions
+3
View File
@@ -2,4 +2,7 @@ from .nodes import *
NODE_CLASS_MAPPINGS = {
"Uncond Zero":uncondZeroNode,
"Conditioning combine positive and negative":cond_combine_pos_neg,
"Conditioning crop or fill":conditioningCropAdd,
"interrupt on NaN": interruptNaNpatch,
}
+185 -58
View File
@@ -1,72 +1,199 @@
import torch
from copy import deepcopy
from comfy.model_management import interrupt_current_processing
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": 1, "min": 0.0, "max": 10.0, "step": 0.01, "round": 0.01}),
"method":(["uncond_zero_v1","uncond_zero_v2","uncond_zero_v3"], {"default": "uncond_zero_v3"},),
"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
def uncond_zero(args):
nonlocal prev_cond
if args["cond_scale"] > 1:
print(f" CFG too high! You may be infering a negative for nothing!")
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(uncond_zero)
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(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
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": 1.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, scale, method):
def uncond_zero(args):
cond = args["cond_denoised"]
x_orig = args["input"]
x_orig -= x_orig.mean() # the main trick to not get a mess is simply to subtract the mean values. I guess SD likes it gaussian AF
cond -= cond.mean()
return x_orig - cond / cond.std() ** .5 * scale # the square root of the std is simply an ever changing scale that fits the bill. The only true condition is to have it not above one near the end.
def uncond_zero_v2(args):
cond = args["cond_denoised"]
x_orig = args["input"]
cond -= cond.mean()
result = torch.zeros_like(x_orig)
for b in range(len(x_orig)):
for c in range(len(x_orig[b])):
x_orig[b][c] -= x_orig[b][c].mean()
cond_c_mean = cond[b][c].mean()
cond[b][c] -= cond_c_mean
result[b][c] = x_orig[b][c] - cond[b][c] / cond[b][c].std() ** .5 * scale + cond_c_mean
return result
# new_scale = 1 / (model.model.latent_format.scale_factor * 8) # Anything below this value gave visible artifacts.
# alright it was around 0.95 with SDXL and 1 is just better. SD 1.x latent scale gives a lower value which ended in bad results.
#Taken and modified from comfy_extras/nodes_model_advanced
def rescale_cfg(args):
x_orig = args["input"]
x_orig -= x_orig.mean()
cond = args["cond_denoised"]
cond -= cond.mean()
cond = x_orig - cond / cond.std() ** .5 * min(scale, 1)
sigma = args["sigma"]
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
#rescale cfg has to be done on v-pred model output
x = x_orig / (sigma * sigma + 1.0)
uncond = x / sigma
cond = ((x - (x_orig - cond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma)
#rescalecfg
x_cfg = uncond + max(scale, 1) * (cond - uncond)
ro_pos = torch.std(cond, dim=(1,2,3), keepdim=True)
ro_cfg = torch.std(x_cfg, dim=(1,2,3), keepdim=True)
x_rescaled = x_cfg * (ro_pos / ro_cfg)
return x_orig - (x - x_rescaled * sigma / (sigma * sigma + 1.0) ** 0.5)
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_cfg_function({"uncond_zero_v1":uncond_zero,"uncond_zero_v2":uncond_zero_v2,"uncond_zero_v3":rescale_cfg}[method])
return (m, )
m.set_model_sampler_post_cfg_function(interrupt_on_nan)
return (m, )