Add files via upload

This commit is contained in:
Extraltodeus
2024-08-12 10:15:44 +02:00
committed by GitHub
parent 2eb785188f
commit 76e6be35e7
2 changed files with 600 additions and 216 deletions
+23 -8
View File
@@ -1,27 +1,42 @@
from .nodes import *
NODE_CLASS_MAPPINGS = {
NODE_CLASS_MAPPINGS = {}
NODE_CLASS_MAPPINGS_ADD = {
"Pre CFG automatic scale": automatic_pre_cfg,
"Pre CFG uncond zero": uncondZeroPreCFGNode,
"Pre CFG perp-neg": pre_cfg_perp_neg,
# "Pre CFG re-negative": pre_cfg_re_negative,
"Pre CFG PAG": perturbed_attention_guidance_pre_cfg_node,
"Pre CFG zero attention": zero_attention_pre_cfg_node,
# "Pre CFG color control": latent_color_control_pre_cfg_node,
"Pre CFG channel multiplier": channel_multiplier_node,
"Pre CFG multiplier": multiply_cond_pre_cfg_node,
"Pre CFG norm neg to pos": norm_uncond_to_cond_pre_cfg_node,
"Pre CFG subtract mean": PreCFGsubtractMeanNode,
"Pre CFG sharpening": condDiffSharpeningNode,
"Pre CFG consensus sharpening": condConsensusSharpeningNode,
"Pre CFG exponentiation": condExpNode,
"Pre CFG variable scaling": variable_scale_pre_cfg_node,
"Pre CFG gradient scaling": gradient_scaling_pre_cfg_node,
"Pre CFG flip flop": flip_flip_conds_pre_cfg_node,
"Pre CFG replace negative channel": replace_uncond_channel_pre_cfg_node,
"Pre CFG merge negative channel": merge_uncond_channel_pre_cfg_node,
"Pre CFG sharpening": condDiffSharpeningNode,
"Pre CFG exponentiation": condExpNode,
"Conditioning set timestep from sigma": ConditioningSetTimestepRangeFromSigma,
"Support empty uncond": support_empty_uncond_pre_cfg_node,
"Shape attention": ShapeAttentionNode,
"Excellent attention": ExlAttentionNode,
"Post CFG subtract mean": PostCFGsubtractMeanNode,
# "Post CFG make a dot": PostCFGDotNode,
"Individual channel selector": individual_channel_selection_node,
"Subtract noise mean": latent_noise_subtract_mean_node,
"Empty RGB image": EmptyRGBImage,
"Gradient RGB image": GradientRGBImage,
}
NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_ADD)
for c in [4,8,16,32,64,128]:
NODE_CLASS_MAPPINGS[f"Channel selector for {c} channels"] = type("channel_selection_node", (channel_selection_node,), { "CHANNELS_AMOUNT": c})
+577 -208
View File
@@ -1,5 +1,6 @@
import torch
from math import ceil
import torch.nn.functional as F
from math import ceil, floor
from copy import deepcopy
import comfy.model_patcher
from comfy.sampler_helpers import convert_cond
@@ -8,15 +9,28 @@ from comfy.ldm.modules.attention import optimized_attention_for_device
from nodes import ConditioningConcat, ConditioningSetTimestepRange
import comfy.model_management as model_management
from comfy.latent_formats import SDXL as SDXL_Latent
import os
from comfy.taesd import taesd as taesd_class
taesd = taesd_class.TAESD()
current_dir = os.path.dirname(os.path.realpath(__file__))
SDXL_Latent = SDXL_Latent()
sdxl_latent_rgb_factors = SDXL_Latent.latent_rgb_factors
ConditioningConcat = ConditioningConcat()
ConditioningSetTimestepRange = ConditioningSetTimestepRange()
default_attention = optimized_attention_for_device(model_management.get_torch_device())
default_device = model_management.get_torch_device()
weighted_average = lambda tensor1, tensor2, weight1: (weight1 * tensor1 + (1 - weight1) * tensor2)
selfnorm = lambda x: x / x.norm()
minmaxnorm = lambda x: torch.nan_to_num((x - x.min()) / (x.max() - x.min()), nan=0.0, posinf=1.0, neginf=0.0)
normlike = lambda x, y: x / x.norm() * y.norm()
def get_sigma_min_max(model):
model_sampling = model.model.model_sampling
sigma_min = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_min)).item()
sigma_max = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)).item()
return sigma_min, sigma_max
class pre_cfg_perp_neg:
@classmethod
@@ -29,6 +43,7 @@ class pre_cfg_perp_neg:
"context_length": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"start_at_sigma": ("FLOAT", {"default": 15, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
# "cond_or_uncond": (["both","uncond"], {"default":"uncond"}),
}
}
RETURN_TYPES = ("MODEL",)
@@ -36,7 +51,7 @@ class pre_cfg_perp_neg:
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, clip, neg_scale, set_context_length, context_length, start_at_sigma, end_at_sigma):
def patch(self, model, clip, neg_scale, set_context_length, context_length, start_at_sigma, end_at_sigma, cond_or_uncond="uncond"):
empty_cond, pooled = clip.encode_from_tokens(clip.tokenize(""), return_pooled=True)
nocond = [[empty_cond, {"pooled_output": pooled}]]
if context_length > 1 and set_context_length:
@@ -69,7 +84,12 @@ class pre_cfg_perp_neg:
perp = neg - ((torch.mul(neg, pos).sum())/(torch.norm(pos)**2)) * pos
perp_neg = perp * neg_scale
conds_out[0] = noise_pred_nocond + pos
if cond_or_uncond == "both":
perp_p = pos - ((torch.mul(neg, pos).sum())/(torch.norm(neg)**2)) * neg
perp_pos = perp_p * neg_scale
conds_out[0] = noise_pred_nocond + perp_pos
else:
conds_out[0] = noise_pred_nocond + pos
conds_out[1] = noise_pred_nocond + perp_neg
return conds_out
@@ -78,6 +98,64 @@ class pre_cfg_perp_neg:
m.set_model_sampler_pre_cfg_function(pre_cfg_perp_neg_function)
return (m, )
class pre_cfg_re_negative:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"clip": ("CLIP",),
"empty_proportion": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 1/20, "round": 0.01}),
"progressive_scale" : ("BOOLEAN", {"default": False}),
"set_context_length" : ("BOOLEAN", {"default": False}),
"context_length": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"end_at_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, clip, empty_proportion, progressive_scale, set_context_length, context_length, end_at_sigma):
sigma_min, sigma_max = get_sigma_min_max(model)
empty_cond, pooled = clip.encode_from_tokens(clip.tokenize(""), return_pooled=True)
nocond = [[empty_cond, {"pooled_output": pooled}]]
if context_length > 1 and set_context_length:
short_nocond = deepcopy(nocond)
for x in range(context_length - 1):
(nocond,) = ConditioningConcat.concat(nocond, short_nocond)
nocond = convert_cond(nocond)
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
sigma = args["sigma"][0]
# cond_scale = args["cond_scale"]
if sigma <= end_at_sigma or not torch.any(conds_out[1]):
return conds_out
model_options = args["model_options"]
timestep = args["timestep"]
model = args["model"]
x_orig = args["input"]
nocond_processed = encode_model_conds(model.extra_conds, nocond, x_orig, x_orig.device, "negative")
(noise_pred_nocond,) = calc_cond_batch(model, [nocond_processed], x_orig, timestep, model_options)
if progressive_scale:
progression = (sigma - sigma_min) / (sigma_max - sigma_min)
current_scale = progression * empty_proportion + (1 - progression) * (1 - empty_proportion)
current_scale = torch.clamp(current_scale, min=0, max=1)
conds_out[1] = current_scale * noise_pred_nocond + conds_out[1] * (1 - current_scale)
else:
conds_out[1] = empty_proportion * noise_pred_nocond + conds_out[1] * (1 - empty_proportion)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
@torch.no_grad()
def normalize_adjust(a,b,strength=1):
norm_a = torch.linalg.norm(a)
@@ -171,6 +249,9 @@ class condExpNode:
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
if do_on in ['both','uncond'] and not uncond:
return conds_out
for b in range(len(conds_out[0])):
if do_on in ['both','cond']:
conds_out[0][b] = normalized_pow(conds_out[0][b], exponent)
@@ -191,11 +272,11 @@ def topk_average(latent, top_k=0.25, measure="average"):
return value_range
apply_scaling_methods = {
"individual": lambda c, m: c * torch.tensor(m).view(c.shape[0],1,1).to(c.device) * 1.25,
"individual": lambda c, m: c * torch.tensor(m).view(c.shape[0],1,1).to(c.device),
"all_as_one": lambda c, m: c * m[0],
"average" : lambda c, m: c * (sum(m) / len(m)),
"smallest": lambda c, m: c * min(m),
"biggest" : lambda c, m: c * max(m),
"average_of_all_channels" : lambda c, m: c * (sum(m) / len(m)),
"smallest_of_all_channels": lambda c, m: c * min(m),
"biggest_of_all_channels" : lambda c, m: c * max(m),
}
measuring_methods = {
@@ -207,28 +288,41 @@ measuring_methods = {
class automatic_pre_cfg:
@classmethod
def INPUT_TYPES(s):
scaling_methods_names = [k for k in apply_scaling_methods]
scaling_methods_names = [k for k in apply_scaling_methods]
measuring_methods_names = [k for k in measuring_methods]
return {"required": {
"model": ("MODEL",),
"scaling_method": (scaling_methods_names, {"default": scaling_methods_names[2]}),
"min_max_method": ([m for m in measuring_methods],),
# "top_k": ("FLOAT", {"default": 0.25, "min": 0.01, "max": 0.5, "step": 1/100, "round": 1/100}),
"scaling_method": (scaling_methods_names, {"default": scaling_methods_names[0]}),
"min_max_method": ([m for m in measuring_methods], {"default": measuring_methods_names[1]}),
"reference_CFG": ("FLOAT", {"default": 8, "min": 0.0, "max": 100, "step": 1/10, "round": 1/100}),
"scale_multiplier": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 100, "step": 1/100, "round": 1/100}),
"top_k": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 0.5, "step": 1/20, "round": 1/100}),
},
"optional": {
"channels_selection": ("CHANS",),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_TYPES = ("MODEL","STRING",)
RETURN_NAMES = ("MODEL","parameters",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, scaling_method, min_max_method="difference", top_k=0.25, channels_selection=None):
def patch(self, model, scaling_method, min_max_method="difference", reference_CFG=8, scale_multiplier=0.8, top_k=0.25, channels_selection=None):
parameters_string = f"scaling_method: {scaling_method}\nmin_max_method: {min_max_method}"
if channels_selection is not None:
for x in range(channels_selection):
parameters_string += f"\nchannel {x+1}: {channels_selection[x]}"
scaling_methods_names = [k for k in apply_scaling_methods]
@torch.no_grad()
def automatic_pre_cfg(args):
conds_out = args["conds_out"]
conds_out = args["conds_out"]
cond_scale = args["cond_scale"]
uncond = torch.any(conds_out[1])
if reference_CFG == 0:
reference_scale = cond_scale
else:
reference_scale = reference_CFG
if not uncond:
return conds_out
@@ -240,24 +334,26 @@ class automatic_pre_cfg:
for b in range(len(conds_out[0])):
chans = []
if scaling_method == scaling_methods_names[1]:
if all(channels):
mes = topk_average(8 * conds_out[0][b] - 7 * conds_out[1][b], top_k=top_k, measure=min_max_method)
mes = topk_average(reference_scale * conds_out[0][b] - (reference_scale - 1) * conds_out[1][b], top_k=top_k, measure=min_max_method)
else:
cond_for_measure = torch.stack([conds_out[0][b][j] for j in range(len(channels)) if channels[j]])
uncond_for_measure = torch.stack([conds_out[1][b][j] for j in range(len(channels)) if channels[j]])
mes = topk_average(8 * cond_for_measure - 7 * uncond_for_measure, top_k=top_k, measure=min_max_method)
chans.append(1 / max(mes,0.01))
mes = topk_average(reference_scale * cond_for_measure - (reference_scale - 1) * uncond_for_measure, top_k=top_k, measure=min_max_method)
chans.append(scale_multiplier / max(mes,0.01))
else:
for c in range(len(conds_out[0][b])):
if not channels[c]:
if scaling_method == scaling_methods_names[0]:
chans.append(1)
continue
mes = topk_average(8 * conds_out[0][b][c] - 7 * conds_out[1][b][c], top_k=top_k, measure=min_max_method)
new_scale = 1 / max(mes,0.01)
mes = topk_average(reference_scale * conds_out[0][b][c] - (reference_scale - 1) * conds_out[1][b][c], top_k=top_k, measure=min_max_method)
new_scale = scale_multiplier / max(mes,0.01)
chans.append(new_scale)
conds_out[0][b] = apply_scaling_methods[scaling_method](conds_out[0][b],chans)
conds_out[1][b] = apply_scaling_methods[scaling_method](conds_out[1][b],chans)
@@ -265,7 +361,7 @@ class automatic_pre_cfg:
m = model.clone()
m.set_model_sampler_pre_cfg_function(automatic_pre_cfg)
return (m, )
return (m, parameters_string,)
class channel_selection_node:
CHANNELS_AMOUNT = 4
@@ -348,7 +444,7 @@ class support_empty_uncond_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"method": (["divide by CFG","from cond"],),
"method": (["from cond","divide by CFG"],),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
@@ -366,7 +462,7 @@ class support_empty_uncond_pre_cfg_node:
if method == "divide by CFG":
conds_out[0] /= cond_scale
else:
conds_out[1] = conds_out[0]
conds_out[1] = conds_out[0].clone()
return conds_out
m = model.clone()
@@ -432,7 +528,8 @@ class zero_attention_pre_cfg_node:
if not cond_generated:
cond_to_process = replace_timestep(cond_to_process)
elif mix_scale == 1:
print(" Mix scale at one!\nPrediction generated for nothing.\nUse the node ConditioningSetTimestepRange to avoid generating if you want to use the full result.")
print(" Mix scale at one!\nPrediction not generated.\nUse the node ConditioningSetTimestepRange to avoid generating if you want to use this node.")
return conds_out
model_options = deepcopy(args["model_options"])
for att in attn:
@@ -500,118 +597,6 @@ class perturbed_attention_guidance_pre_cfg_node:
m.set_model_sampler_pre_cfg_function(perturbed_attention_guidance_pre_cfg_patch)
return (m, )
@torch.no_grad()
def euclidean_weights(tensors,exponent=2,proximity_exponent=1,min_score_into_zeros=0):
divider = tensors.shape[0]
if exponent == 0:
exponent = 6.55
device = tensors.device
distance_weights = torch.zeros_like(tensors).to(device = device)
for i in range(len(tensors)):
for j in range(len(tensors)):
if i == j: continue
current_distance = (tensors[i] - tensors[j]).abs() / divider
if proximity_exponent > 1:
current_distance = current_distance ** proximity_exponent
distance_weights[i] += current_distance
min_stack, _ = torch.min(distance_weights, dim=0)
max_stack, _ = torch.max(distance_weights, dim=0)
max_stack = torch.where(max_stack == 0, torch.tensor(1), max_stack)
sum_of_weights = torch.zeros_like(tensors[0]).to(device = device)
max_stack -= min_stack
for i in range(len(tensors)):
distance_weights[i] -= min_stack
distance_weights[i] /= max_stack
distance_weights[i] = 1 - distance_weights[i]
distance_weights[i] = torch.clamp(distance_weights[i], min=0)
if min_score_into_zeros > 0:
distance_weights[i] = torch.where(distance_weights[i] < min_score_into_zeros, torch.zeros_like(distance_weights[i]), distance_weights[i])
if exponent > 1:
distance_weights[i] = distance_weights[i] ** exponent
sum_of_weights += distance_weights[i]
mean_score = (sum_of_weights.mean() / divider) ** exponent
sum_of_weights = torch.where(sum_of_weights == 0, torch.zeros_like(sum_of_weights) + 1 / divider, sum_of_weights)
result = torch.zeros_like(tensors[0]).to(device = device)
for i in range(len(tensors)):
distance_weights[i] /= sum_of_weights
distance_weights[i] = torch.where(torch.isnan(distance_weights[i]) | torch.isinf(distance_weights[i]), torch.zeros_like(distance_weights[i]), distance_weights[i])
result = result + tensors[i] * distance_weights[i]
return result, mean_score
class condConsensusSharpeningNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"scale": ("FLOAT", {"default": 0.75, "min": -10.0, "max": 10.0, "step": 1/20, "round": 1/100}),
"start_at_sigma": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, scale, start_at_sigma, end_at_sigma):
model_sampling = model.model.model_sampling
sigma_max = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)).item()
prev_conds = []
prev_unconds = []
@torch.no_grad()
def sharpen_conds_pre_cfg(args):
nonlocal prev_conds, prev_unconds
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
sigma = args["sigma"][0].item()
if sigma <= end_at_sigma:
return conds_out
first_step = sigma > (sigma_max - 1)
if first_step:
prev_conds = []
prev_unconds = []
prev_conds.append(conds_out[0] / conds_out[0].norm())
if uncond:
prev_unconds.append(conds_out[1] / conds_out[1].norm())
if sigma > start_at_sigma:
return conds_out
if len(prev_conds) > 3:
consensus_cond, mean_score_cond = euclidean_weights(torch.stack(prev_conds))
consensus_cond = consensus_cond * conds_out[0].norm()
if len(prev_unconds) > 3 and uncond:
consensus_uncond, mean_score_uncond = euclidean_weights(torch.stack(prev_unconds))
consensus_uncond = consensus_uncond * conds_out[1].norm()
for b in range(len(conds_out[0])):
for c in range(len(conds_out[0][b])):
if len(prev_conds) > 3:
conds_out[0][b][c] = normalize_adjust(conds_out[0][b][c], consensus_cond[b][c], mean_score_cond * scale)
if len(prev_unconds) > 3 and uncond:
conds_out[1][b][c] = normalize_adjust(conds_out[1][b][c], consensus_uncond[b][c], mean_score_uncond * scale)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(sharpen_conds_pre_cfg)
return (m, )
def sigma_to_percent(model_sampling, sigma_value):
if sigma_value >= 999999999.9:
return 0.0
@@ -661,7 +646,7 @@ class ShapeAttentionNode:
CATEGORY = "model_patches"
def patch(self, model, scale, start_at_sigma=999999999.9, end_at_sigma=0.0, enabled=True, attention="both", unet_block="input", unet_block_id=8):
def patch(self, model, scale, start_at_sigma=999999999.9, end_at_sigma=0.0, enabled=True, attention="self", unet_block="input", unet_block_id=8):
attn = {"both":["attn1","attn2"],"self":["attn1"],"cross":["attn2"]}[attention]
if scale == 1:
print(" Shape attention disabled (scale is one)")
@@ -684,13 +669,40 @@ class ShapeAttentionNode:
return (m,)
class ExlAttentionNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"scale": ("FLOAT", {"default": 2, "min": -1.0, "max": 10.0, "step": 1/10, "round": 1/100}),
"enabled": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches"
def patch(self, model, scale, enabled):
if not enabled:
return (model,)
m = model.clone()
def cross_patch(q, k, v, extra_options, mask=None):
first_attention = default_attention(q, k, v, extra_options['n_heads'], mask)
second_attention = normlike(q+(q-default_attention(first_attention, k, v, extra_options['n_heads'])), first_attention) * scale
return second_attention
m.model_options = comfy.model_patcher.set_model_options_patch_replace(m.model_options, cross_patch, "attn2", "middle", 0)
return (m,)
class PreCFGsubtractMeanNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
# "per_channel" : ("BOOLEAN", {"default": False}),
# "per_channel" : ("BOOLEAN", {"default": False}), #It's just not good
"start_at_sigma": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"enabled" : ("BOOLEAN", {"default": True}),
@@ -728,7 +740,7 @@ class PostCFGsubtractMeanNode:
return {
"required": {
"model": ("MODEL",),
# "per_channel" : ("BOOLEAN", {"default": False}),
# "per_channel" : ("BOOLEAN", {"default": False}), #It's just not good
"start_at_sigma": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"enabled" : ("BOOLEAN", {"default": True}),
@@ -802,9 +814,13 @@ class PostCFGDotNode:
class uncondZeroPreCFGNode:
@classmethod
def INPUT_TYPES(s):
scaling_methods_names = [k for k in apply_scaling_methods]
return {"required": {
"model": ("MODEL",),
"scale": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 10.0, "step": 1/20, "round": 0.01}),
"scale": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 10.0, "step": 1/20, "round": 0.01}),
"start_at_sigma": ("FLOAT", {"default": 100, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
"scaling_method": (scaling_methods_names, {"default": scaling_methods_names[2]}),
}
}
RETURN_TYPES = ("MODEL",)
@@ -812,85 +828,36 @@ class uncondZeroPreCFGNode:
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, scale):
model_sampling = model.model.model_sampling
sigma_max = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)).item()
def patch(self, model, scale, start_at_sigma, end_at_sigma, scaling_method):
scaling_methods_names = [k for k in apply_scaling_methods]
@torch.no_grad()
def uncond_zero_pre_cfg(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
sigma = args["sigma"][0].item()
if uncond or sigma < (sigma_max * 0.069):
if uncond or sigma <= end_at_sigma or sigma > start_at_sigma:
return conds_out
for b in range(len(conds_out[0])):
chans = []
if scaling_method == scaling_methods_names[1]:
mes = topk_average(8 * conds_out[0][b] - 7 * conds_out[1][b], measure="difference")
for c in range(len(conds_out[0][b])):
mes = topk_average(conds_out[0][b][c], measure="difference") ** 0.5
conds_out[0][b][c] = conds_out[0][b][c] * scale / mes
chans.append(scale / mes)
conds_out[0][b] = apply_scaling_methods[scaling_method](conds_out[0][b],chans)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(uncond_zero_pre_cfg)
return (m, )
class latent_color_control_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"Red": ("FLOAT", {"default": 0, "min": -2.0, "max": 2.0, "step": 1/100, "round": 1/100}),
"Green": ("FLOAT", {"default": 0, "min": -2.0, "max": 2.0, "step": 1/100, "round": 1/100}),
"Blue": ("FLOAT", {"default": 0, "min": -2.0, "max": 2.0, "step": 1/100, "round": 1/100}),
"selection": (["both","cond","uncond"],),
"start_at_sigma": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 100.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 01.0, "min": 0.0, "max": 100.0, "step": 1/100, "round": 1/100}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, Red, Green, Blue, selection, start_at_sigma, end_at_sigma):
latent_rgb_factors = sdxl_latent_rgb_factors
rgb = [Red, Green, Blue]
@torch.no_grad()
def latent_control_pre_cfg_function(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
sigma = args["sigma"]
if sigma[0] <= end_at_sigma or sigma[0] > start_at_sigma or all(c == 0 for c in rgb):
return conds_out
conds_index = []
if selection in ["both","cond"]:
conds_index.append(0)
if uncond and selection in ["both","uncond"]:
conds_index.append(1)
for i in conds_index:
for b in range(len(conds_out[i])):
cond_norm = conds_out[i][b].norm()
color_cond = torch.zeros_like(conds_out[i][b])
for c in range(len(conds_out[i][b])):
for r in range(len(rgb)):
if rgb[r] != 0:
color_cond[c] += rgb[r] * latent_rgb_factors[c][r] / 0.13025
conds_out[i][b] = selfnorm(conds_out[i][b] + color_cond) * cond_norm
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(latent_control_pre_cfg_function)
return (m, )
class variable_scale_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"start_multiplier": ("FLOAT", {"default": 1.5, "min": 0.0, "max": 10.0, "step": 1/100, "round": 1/100}),
"end_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 1/100, "round": 1/100}),
"target_scale": ("FLOAT", {"default": 5.0, "min": 1.0, "max": 100.0, "step": 1/2, "round": 1/100}),
"target_as_start": ("BOOLEAN", {"default": True}),
"proportional_to": (["sigma","steps progression"],),
}}
RETURN_TYPES = ("MODEL",)
@@ -898,32 +865,434 @@ class variable_scale_pre_cfg_node:
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, start_multiplier, end_multiplier, proportional_to):
def patch(self, model, target_scale, target_as_start, proportional_to):
model_sampling = model.model.model_sampling
sigma_max = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)).item()
@torch.no_grad()
def variable_scale_pre_cfg_patch(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
conds_out = args["conds_out"]
cond_scale = args["cond_scale"]
sigma = args["sigma"][0].item()
scales = [cond_scale,target_scale]
if not uncond:
if not torch.any(conds_out[1]):
return conds_out
sigma = args["sigma"][0].item()
if proportional_to == "steps progression":
progression = sigma_to_percent(model_sampling, sigma)
else:
progression = 1 - sigma / sigma_max
progression = max(min(progression, 1), 0)
current_multiplier = start_multiplier * (1 - progression) + end_multiplier * progression
conds_out[0] = conds_out[0] * current_multiplier
conds_out[1] = conds_out[1] * current_multiplier
current_scale = scales[target_as_start] * (1 - progression) + scales[not target_as_start] * progression
new_scale = (current_scale - 1) / (cond_scale - 1)
conds_out[1] = weighted_average(conds_out[1], conds_out[0], new_scale)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(variable_scale_pre_cfg_patch)
return (m, )
return (m, )
class latent_noise_subtract_mean_node:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {"required": {
"latent_input": ("LATENT", {"forceInput": True}),
"enabled" : ("BOOLEAN", {"default": True}),
}}
FUNCTION = "exec"
RETURN_TYPES = ("LATENT",)
CATEGORY = "latent"
def exec(self, latent_input, enabled):
if not enabled:
return (latent_input,)
new_latents = deepcopy(latent_input)
for x in range(len(new_latents['samples'])):
new_latents['samples'][x] -= torch.mean(new_latents['samples'][x])
return (new_latents,)
class flip_flip_conds_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"enabled" : ("BOOLEAN", {"default": True})
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, enabled):
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
if not uncond or not enabled:
return conds_out
conds_out[0], conds_out[1] = conds_out[1], conds_out[0]
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
class norm_uncond_to_cond_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"enabled" : ("BOOLEAN", {"default": True})
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, enabled):
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
if not uncond or not enabled:
return conds_out
conds_out[1] = conds_out[1] / conds_out[1].norm() * conds_out[0].norm()
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
class replace_uncond_channel_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"channel": ("INT", {"default": 1, "min": 1, "max": 128, "step": 1}),
"enabled" : ("BOOLEAN", {"default": True})
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, channel, enabled):
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
if not uncond or not enabled:
return conds_out
for b in range(len(conds_out[0])):
if len(conds_out[1][b]) < channel:
print(F" WRONG CHANNEL SELECTED. THE LATENT SPACE ONLY HAS {len(conds_out[1][b])} CHANNELS")
else:
conds_out[1][b][channel - 1] = conds_out[0][b][channel - 1]
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
class merge_uncond_channel_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"channel": ("INT", {"default": 1, "min": 1, "max": 128, "step": 1}),
"CFG_scale": ("FLOAT", {"default": 5, "min": 2.0, "max": 100.0, "step": 1/2, "round": 1/100}),
"start_at_sigma": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 100.0, "step": 1/100, "round": 1/100}),
"end_at_sigma": ("FLOAT", {"default": 01.0, "min": 0.0, "max": 100.0, "step": 1/100, "round": 1/100}),
"enabled" : ("BOOLEAN", {"default": True})
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, channel, CFG_scale, start_at_sigma, end_at_sigma, enabled):
if not enabled: return model,
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
cond_scale = args["cond_scale"]
sigma = args["sigma"][0].item()
if not torch.any(conds_out[1]) or sigma <= end_at_sigma or sigma > start_at_sigma:
return conds_out
for b in range(len(conds_out[0])):
if len(conds_out[1][b]) < channel:
print(F" WRONG CHANNEL SELECTED. THE LATENT SPACE ONLY HAS {len(conds_out[1][b])} CHANNELS")
else:
new_scale = (CFG_scale - 1) / (cond_scale - 1)
conds_out[1][b][channel - 1] = weighted_average(conds_out[1][b][channel - 1], conds_out[0][b][channel - 1], new_scale)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
class multiply_cond_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"selection": (["both","cond","uncond"],),
"value": ("FLOAT", {"default": 0, "min": -100.0, "max": 100.0, "step": 1/100, "round": 1/100}),
"enabled" : ("BOOLEAN", {"default": True})
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
def patch(self, model, selection, value, enabled):
@torch.no_grad()
def pre_cfg_patch(args):
conds_out = args["conds_out"]
uncond = torch.any(conds_out[1])
if (not uncond and selection in ["both","uncond"]) or not enabled:
return conds_out
if selection in ["both","cond"]:
conds_out[0] = conds_out[0] * value
if selection in ["both","uncond"]:
conds_out[1] = conds_out[1] * value
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
def generate_gradient_mask(tensor, horizontal=False):
dim = 3 if horizontal else 2
gradient = torch.linspace(0, 1, steps=tensor.size(dim), device=tensor.device)
if horizontal:
merging_gradient = gradient.repeat(tensor.size(0), tensor.size(1), tensor.size(2), 1)
else:
merging_gradient = gradient.unsqueeze(1).repeat(tensor.size(0), tensor.size(1), 1, tensor.size(3))
return merging_gradient
class gradient_scaling_pre_cfg_node:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"maximum_scale": ("FLOAT", {"default": 80, "min": 3.0, "max": 1000.0, "step": 1, "round": 1/100, "tooltip":"It is an equivalent to the CFG scale."}),
"minimum_scale": ("FLOAT", {"default": 4.5, "min": 1.0, "max": 10.0, "step": 1/2, "round": 1/100, "tooltip":"It is an equivalent to the CFG scale."}),
"strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 10.0, "step": 1/10, "round": 1/10}),
"end_at_sigma": ("FLOAT", {"default": 0.28, "min": 0.0, "max": 1000.0, "step": 1/100, "round": 1/100}),
# "free_scale" : ("BOOLEAN", {"default": False}),
"converging_scales" : ("BOOLEAN", {"default": True}),
# "noise_add_diff" : ("BOOLEAN", {"default": True}),
# "split_channels" : ("BOOLEAN", {"default": False}),
"invert_mask" : ("BOOLEAN", {"default": False}),
},
"optional":{
"input_mask": ("MASK", {"tooltip":"If only a mask is connected the scale becomes a CFG scale of what is being masked.\nWhen a latent is connected the mask defines what will be modified by the node."},),
"input_latent": ("LATENT", {"tooltip":"If a latent is connected the scale becomes the maximum scale allowed in which to seek similarity."},),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model_patches/Pre CFG"
@torch.no_grad()
def make_new_uncond_at_scale(self,cond,uncond,cond_scale,new_scale):
new_scale_ratio = (new_scale - 1) / (cond_scale - 1)
return cond * (1 - new_scale_ratio) + uncond * new_scale_ratio
@torch.no_grad()
def get_denoised_at_scale(self,x_orig,cond,uncond,cond_scale):
return x_orig - ((x_orig - uncond) + cond_scale * ((x_orig - cond) - (x_orig - uncond)))
def get_latent_guidance_mask_channel(self,x_orig,cond,uncond,guide,minimum_scale,maximum_scale,noise_add_diff):
scales = torch.zeros_like(x_orig, device=x_orig.device)
for b in range(cond.shape[0]):
for c in range(cond.shape[1]):
scales[b][c] = self.get_latent_guidance_mask(x_orig[b][c],cond[b][c],uncond[b][c],guide[0][c],minimum_scale,maximum_scale,noise_add_diff)
return scales
@torch.no_grad()
def get_latent_guidance_mask(self,x_orig,cond,uncond,guide,minimum_scale,maximum_scale,noise_add_diff):
low_denoised = self.get_denoised_at_scale(x_orig,cond,uncond,minimum_scale)
high_denoised = self.get_denoised_at_scale(x_orig,cond,uncond,maximum_scale)
if noise_add_diff:
guide = guide + (guide - (x_orig * guide.norm() / x_orig.norm()))
guide = guide / guide.norm()
low_diff = (low_denoised - guide * low_denoised.norm()).abs()
high_diff = (high_denoised - guide * high_denoised.norm()).abs()
return torch.clamp(low_diff / high_diff, min=0, max=1)
def patch(self, model, maximum_scale, minimum_scale, invert_mask, strength, end_at_sigma, noise_add_diff=True, converging_scales=False, split_channels=False, free_scale=False, input_mask=None, input_latent=None):
if input_mask is None and input_latent is None:
return (model,)
sigma_min, sigma_max = get_sigma_min_max(model)
model_sampling = model.model.model_sampling
scaling_function = self.get_latent_guidance_mask_channel if split_channels else self.get_latent_guidance_mask
mask_as_weight = None
latent_as_guidance = None
if input_mask is not None:
mask_as_weight = input_mask.clone().to(device=default_device)
if invert_mask:
mask_as_weight = 1 - mask_as_weight
if mask_as_weight.dim() == 3:
mask_as_weight = mask_as_weight.unsqueeze(1)
if input_latent is not None:
latent_as_guidance = input_latent["samples"].clone().to(device=default_device)
@torch.no_grad()
def pre_cfg_patch(args):
nonlocal mask_as_weight, latent_as_guidance
conds_out = args["conds_out"]
cond_scale = args["cond_scale"]
x_orig = args['input']
sigma = args["sigma"][0]
sp = min(1,max(0,sigma_to_percent(model_sampling, sigma - sigma_min * 3) + 1 / 100)) ** 2
if not torch.any(conds_out[1]) or sigma <= end_at_sigma or (converging_scales and sp == 1):
return conds_out
if converging_scales:
current_maximum_scale = sp * cond_scale + (1 - sp) * maximum_scale
current_minimum_scale = sp * cond_scale + (1 - sp) * minimum_scale
else:
current_maximum_scale = maximum_scale
current_minimum_scale = minimum_scale
if mask_as_weight is not None and mask_as_weight.shape[-2:] != conds_out[1].shape[-2:]:
mask_as_weight = F.interpolate(mask_as_weight, size=(conds_out[1].shape[-2], conds_out[1].shape[-1]), mode='bilinear', align_corners=False)
if latent_as_guidance is not None:
if latent_as_guidance.shape[-2:] != conds_out[1].shape[-2:]:
latent_as_guidance = F.interpolate(latent_as_guidance, size=(conds_out[1].shape[-2], conds_out[1].shape[-1]), mode='bilinear', align_corners=False)
scaling_weight = scaling_function(x_orig,conds_out[0],conds_out[1],latent_as_guidance.clone(),current_minimum_scale,current_maximum_scale,noise_add_diff)
target_scales = scaling_weight * current_maximum_scale + (1 - scaling_weight) * current_minimum_scale
if free_scale:
target_scales = target_scales * cond_scale / target_scales.mean()
global_multiplier = strength
if mask_as_weight is not None:
global_multiplier = global_multiplier * mask_as_weight
target_scales = target_scales * global_multiplier + torch.full_like(target_scales, cond_scale) * (1 - global_multiplier)
conds_out[1] = self.make_new_uncond_at_scale(conds_out[0],conds_out[1],cond_scale,target_scales)
return conds_out
else:
target_scales = maximum_scale * mask_as_weight * strength + torch.full_like(conds_out[1], cond_scale) * (1 - mask_as_weight * strength)
conds_out[1] = self.make_new_uncond_at_scale(conds_out[0],conds_out[1],cond_scale,maximum_scale)
return conds_out
m = model.clone()
m.set_model_sampler_pre_cfg_function(pre_cfg_patch)
return (m, )
class EmptyRGBImage:
@classmethod
def INPUT_TYPES(s):
return {"required": { "width": ("INT", {"default": 1024, "min": 1, "max": 16384, "step": 1}),
"height": ("INT", {"default": 1024, "min": 1, "max": 16384, "step": 1}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096}),
"r": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
"g": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
"b": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
},
"optional": {
"grayscale_to_color": ("IMAGE",),
}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "generate"
CATEGORY = "image"
def generate(self, width, height, batch_size=1, r=0, g=0, b=0, grayscale_to_color=None):
if grayscale_to_color is not None:
grayscale_to_color = grayscale_to_color.permute(0, 3, 1, 2).mean(dim=1).unsqueeze(-1)
height = grayscale_to_color.shape[1]
width = grayscale_to_color.shape[2]
r_normalized = torch.full([batch_size, height, width, 1], r / 255.0)
g_normalized = torch.full([batch_size, height, width, 1], g / 255.0)
b_normalized = torch.full([batch_size, height, width, 1], b / 255.0)
rgb_image = torch.cat((r_normalized, g_normalized, b_normalized), dim=-1)
if grayscale_to_color is not None:
rgb_image = rgb_image * grayscale_to_color
return (rgb_image,)
gradient_patterns = {
"linear": lambda x, y: x,
"sine": lambda x, y: torch.sin(x * torch.pi * y),
"triangle": lambda x, y: 2 * torch.abs(torch.round(x % (1 / max(y, 1)) * y) - (x % (1 / max(y, 1)) * y)),
}
class GradientRGBImage:
@classmethod
def INPUT_TYPES(s):
return {"required": { "width": ("INT", {"default": 1024, "min": 0, "max": 16384, "step": 64}),
"height": ("INT", {"default": 1024, "min": 0, "max": 16384, "step": 64}),
"r1": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
"g1": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
"b1": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
"r2": ("INT", {"default": 255, "min": 0, "max": 255, "step": 1}),
"g2": ("INT", {"default": 255, "min": 0, "max": 255, "step": 1}),
"b2": ("INT", {"default": 255, "min": 0, "max": 255, "step": 1}),
"axis" : (["vertical","horizontal","circular"],),
"power_to": ("INT", {"default": 1, "min": 1, "max": 16, "step": 1}),
"reverse_power" : ("BOOLEAN", {"default": False}),
},
"optional":{
"mask": ("MASK",),
}
}
RETURN_TYPES = ("IMAGE","MASK",)
FUNCTION = "generate"
CATEGORY = "image"
def get_gradient_mask(self,width,height,horizontal):
if horizontal:
return torch.linspace(0, 1, width).view(1, 1, width).repeat(1, height, 1)
return torch.linspace(0, 1, height).view(1, height, 1).repeat(1, 1, width)
def generate(self, width, height, batch_size=1, r1=0, g1=0, b1=0, r2=255, g2=255, b2=255, pattern_value=1, power_to=1, reverse_power=False, axis="vertical", mask=None):
gradient = self.get_gradient_mask(width, height, axis in ["horizontal","circular"])
gradient = gradient_patterns["linear" if axis != "circular" else "sine"](gradient, pattern_value)
if axis == "circular":
gradient2 = self.get_gradient_mask(width, height, False)
gradient2 = gradient_patterns["sine"](gradient2, pattern_value)
gradient = gradient * gradient2
if power_to > 1:
if reverse_power: gradient = 1 - gradient
gradient = gradient ** power_to
if reverse_power: gradient = 1 - gradient
if mask is not None:
if mask.shape != gradient.shape:
mask = F.interpolate(mask.unsqueeze(1), size=(gradient.shape[-2], gradient.shape[-1]), mode='nearest').squeeze(1)
gradient = gradient * mask
gradient = gradient.squeeze(0).unsqueeze(-1)
r_gradient = r1 / 255.0 + gradient * (r2 - r1) / 255.0
g_gradient = g1 / 255.0 + gradient * (g2 - g1) / 255.0
b_gradient = b1 / 255.0 + gradient * (b2 - b1) / 255.0
r_image = r_gradient.expand(batch_size, height, width, 1)
g_image = g_gradient.expand(batch_size, height, width, 1)
b_image = b_gradient.expand(batch_size, height, width, 1)
rgb_image = torch.cat((r_image, g_image, b_image), dim=-1)
mask_gradient = gradient.expand(1, height, width, 1).squeeze(-1)
return (rgb_image,mask_gradient,)