diff --git a/__init__.py b/__init__.py index 1c24038..0a61239 100644 --- a/__init__.py +++ b/__init__.py @@ -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}) diff --git a/nodes.py b/nodes.py index 8f04aee..11327c7 100644 --- a/nodes.py +++ b/nodes.py @@ -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, ) \ No newline at end of file + 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,) \ No newline at end of file