From 5570e3c3fde2c317cbb4a1fecb94bc4778b03466 Mon Sep 17 00:00:00 2001 From: peter942 Date: Thu, 7 Dec 2023 16:16:13 +0100 Subject: [PATCH] Refactoring --- .DS_Store | Bin 6148 -> 0 bytes .gitignore | 6 + control/nodes.py => SteerableMotion.py | 347 ++---------- __init__.py | 2 +- control/.DS_Store | Bin 6148 -> 0 bytes control/deprecated_nodes.py | 70 --- control/latent_keyframe_nodes.py | 296 ----------- control/logger.py | 36 -- control/resampler.py | 121 ----- control/weight_nodes.py | 157 ------ .../AdvancedControlNet.py | 498 +++++++++++++++--- {control => imports}/IPAdapterPlus.py | 305 +++++------ 12 files changed, 618 insertions(+), 1220 deletions(-) delete mode 100644 .DS_Store rename control/nodes.py => SteerableMotion.py (57%) delete mode 100644 control/.DS_Store delete mode 100644 control/deprecated_nodes.py delete mode 100644 control/latent_keyframe_nodes.py delete mode 100644 control/logger.py delete mode 100644 control/resampler.py delete mode 100644 control/weight_nodes.py rename control/control.py => imports/AdvancedControlNet.py (51%) rename {control => imports}/IPAdapterPlus.py (77%) diff --git a/.DS_Store b/.DS_Store deleted file mode 100644 index 0d2bc53b7cab718738156292dde2564c16f131d3..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 6148 zcmeHKKX21O6n~c(a%lz1K&7%-u%U{^9S{|Z3F*LwFrovXVAoM8UVmi37sWu zP7(5xE=WaNt|BseM%wd&9g+1HBCWwPU>UfG4A8ghz-{P30)_egR)4&?Vfr%BZbzwz z(8k|=_`d#e@6-KxfBaW}fA?f>sQ!F}sZfI*5RgLw1H{yv9pKCUVuV@4_N~uvmsAaZ zJSP(;m7E=A)x?&)P$CJtni2qlcmC zGdcZs{TH9TSIvdV=4pPD$I6l1`e!`84hDHD@*d82T<5Gn+c|dz^gH3p=U#~|z%pPN zxM&Q}{@|h#`Wj1xa_hiGUI7ptG)uug-6bf;)#z(16`}`)sZ>Oj%Jdb3sdTivI?mTv zDpcvfW$Y7+_92 zXt!`l`fOcT9G$f;>Pu7-iYpb$6m0ZyEIV`*ucAsppGzG?Ut_5dEhzR!K+<3v%fLTn F;5W>{rcM9= diff --git a/.gitignore b/.gitignore index 68bc17f..d60ebac 100644 --- a/.gitignore +++ b/.gitignore @@ -82,6 +82,12 @@ target/ profile_default/ ipython_config.py + +.DS_Store + +# Ignore Python cache files + + # pyenv # For a library or package, you might want to ignore these files since the code is # intended to run in multiple environments; otherwise, check them in: diff --git a/control/nodes.py b/SteerableMotion.py similarity index 57% rename from control/nodes.py rename to SteerableMotion.py index 5e1f31e..04d7a2c 100644 --- a/control/nodes.py +++ b/SteerableMotion.py @@ -1,202 +1,22 @@ -import numpy as np -import torch -import folder_paths -from PIL import Image from ast import literal_eval -from .control import ControlNetAdvancedImport, T2IAdapterAdvancedImport, load_controlnet, ControlNetWeightsTypeImport, T2IAdapterWeightsTypeImport,\ - LatentKeyframeGroupImport, TimestepKeyframeImport, TimestepKeyframeGroupImport, is_advanced_controlnet -import matplotlib.pyplot as plt -from .IPAdapterPlus import contrast_adaptive_sharpening, IPAdapterApply,prep_image - -from .weight_nodes import ScaledSoftControlNetWeightsImport, SoftControlNetWeightsImport, CustomControlNetWeightsImport, \ - SoftT2IAdapterWeightsImport, CustomT2IAdapterWeightsImport -from .latent_keyframe_nodes import LatentKeyframeGroupNodeImport, LatentKeyframeInterpolationNodeImport, LatentKeyframeBatchedGroupNodeImport, LatentKeyframeNodeImport,calculate_weights -from .deprecated_nodes import LoadImagesFromDirectory -from .logger import logger -import torchvision.transforms as TT -import torch.nn.functional as F - -import comfy.utils -import comfy.model_management -from comfy.clip_vision import clip_preprocess -from comfy.ldm.modules.attention import optimized_attention -# import BytesIO from io import BytesIO +import torch +import torchvision.transforms as TT +from PIL import Image +import matplotlib.pyplot as plt +import folder_paths - -class TimestepKeyframeNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), - }, - "optional": { - "control_net_weights": ("CONTROL_NET_WEIGHTS", ), - "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), - "latent_keyframe": ("LATENT_KEYFRAME", ), - "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), - } - } +from .imports.IPAdapterPlus import IPAdapterApplyImport, prep_image +from .imports.AdvancedControlNet import ( + calculate_weights, + LatentKeyframeInterpolationNodeImport, + ScaledSoftControlNetWeightsImport, + ControlNetLoaderAdvancedImport, + AdvancedControlNetApplyImport, + TimestepKeyframeNodeImport, +) - RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) - FUNCTION = "load_keyframe" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, - start_percent: float, - control_net_weights: ControlNetWeightsTypeImport=None, - t2i_adapter_weights: T2IAdapterWeightsTypeImport=None, - latent_keyframe: LatentKeyframeGroupImport=None, - prev_timestep_keyframe: TimestepKeyframeGroupImport=None): - if not prev_timestep_keyframe: - prev_timestep_keyframe = TimestepKeyframeGroupImport() - keyframe = TimestepKeyframeImport(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe) - prev_timestep_keyframe.add(keyframe) - return (prev_timestep_keyframe,) - - -class ControlNetLoaderAdvancedImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "control_net_name": (folder_paths.get_filename_list("controlnet"), ), - }, - "optional": { - "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), - } - } - - RETURN_TYPES = ("CONTROL_NET", ) - FUNCTION = "load_controlnet" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" - - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroupImport=None): - controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - controlnet = load_controlnet(controlnet_path, timestep_keyframe) - return (controlnet,) - - -class DiffControlNetLoaderAdvancedImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "control_net_name": (folder_paths.get_filename_list("controlnet"), ) - }, - "optional": { - "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), - } - } - - RETURN_TYPES = ("CONTROL_NET", ) - FUNCTION = "load_controlnet" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" - - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroupImport, model): - controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) - return (controlnet,) - - -class AdvancedControlNetApplyImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "control_net": ("CONTROL_NET", ), - "image": ("IMAGE", ), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), - "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) - }, - "optional": { - "mask_optional": ("MASK", ), - } - } - - RETURN_TYPES = ("CONDITIONING","CONDITIONING") - RETURN_NAMES = ("positive", "negative") - FUNCTION = "apply_controlnet" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" - - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=None): - if strength == 0: - return (positive, negative) - - control_hint = image.movedim(-1,1) - cnets = {} - - out = [] - for conditioning in [positive, negative]: - c = [] - - for t in conditioning: - d = t[1].copy() - - prev_cnet = d.get('control', None) - if prev_cnet in cnets: - c_net = cnets[prev_cnet] - - else: - c_net = control_net.copy().set_cond_hint(control_hint, strength, (start_percent, end_percent)) - # set cond hint mask - if mask_optional is not None: - if is_advanced_controlnet(c_net): - # if not in the form of a batch, make it so - if len(mask_optional.shape) < 3: - mask_optional = mask_optional.unsqueeze(0) - c_net.set_cond_hint_mask(mask_optional) - c_net.set_previous_controlnet(prev_cnet) - cnets[prev_cnet] = c_net - - d['control'] = c_net - d['control_apply_to_uncond'] = False - n = [t[0], d] - c.append(n) - out.append(c) - return (out[0], out[1]) - - -class MaskGeneratorNode: - - RETURN_TYPES = ("MASK",) - FUNCTION = "generate_masks" - CATEGORY = "Steerable-Motion/Interpolation" - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "number_of_masks": ("INT", {"default": 16, "min": 1, "max": 100, "step": 1}), - "strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), - "width": ("INT", {"default": 512, "min": 16, "max": 4096, "step": 1}), - "height": ("INT", {"default": 512, "min": 16, "max": 4096, "step": 1}), - }, - } - - def generate_masks(self, number_of_masks, strength, width, height): - - masks = [] - for _ in range(number_of_masks): - mask = torch.full((height, width), strength) - masks.append(mask) - - # Convert list of masks to a single tensor - masks_tensor = torch.stack(masks, dim=0) - return masks_tensor - - - class BatchCreativeInterpolationNode: @classmethod def IS_CHANGED(cls, **kwargs): @@ -222,8 +42,7 @@ class BatchCreativeInterpolationNode: "type_of_cn_strength_distribution": (["linear", "dynamic"],), "linear_cn_strength_value": ("STRING", {"multiline": False, "default": "(0.0,0.4)"}), "dynamic_cn_strength_values": ("STRING", {"multiline": True, "default": "(0.0,1.0),(0.0,1.0),(0.0,1.0),(0.0,1.0)"}), - "soft_scaled_cn_weights_multiplier": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 10.0, "step": 0.1}), - # "interpolation": (["ease-in-out", "ease-in", "ease-out"],), + "soft_scaled_cn_weights_multiplier": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 10.0, "step": 0.1}), "buffer": ("INT", {"default": 4, "min": 0, "max": 16, "step": 1}), "relative_ipadapter_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}), "relative_ipadapter_influence": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}), @@ -325,9 +144,6 @@ class BatchCreativeInterpolationNode: # Hardcoded dimensions width, height = 512, 512 - # Calculate the reversed weights in a generalizable way (e.g., 0.6 becomes 0.4, 0.1 becomes 0.9) - reversed_weights = [1.0 - weight for weight in weights] - # Map frames to their corresponding reversed weights for easy lookup frame_to_weight = {frame: weights[i] for i, frame in enumerate(frames)} @@ -439,143 +255,92 @@ class BatchCreativeInterpolationNode: keyframe_positions = get_keyframe_positions(type_of_frame_distribution, dynamic_frame_distribution_values, images, linear_frame_distribution_value) cn_strength_values = extract_start_and_endpoint_values(type_of_cn_strength_distribution, dynamic_cn_strength_values, keyframe_positions, linear_cn_strength_value) key_frame_influence_values = extract_keyframe_values(type_of_key_frame_influence, dynamic_key_frame_influence_values, keyframe_positions, linear_key_frame_influence_value) - influence_ranges = calculate_dynamic_influence_ranges(keyframe_positions,key_frame_influence_values) - + influence_ranges = calculate_dynamic_influence_ranges(keyframe_positions,key_frame_influence_values) influence_ranges = add_starting_buffer(influence_ranges, buffer) cn_strength_values = [literal_eval(val) if isinstance(val, str) else val for val in cn_strength_values] + cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights = [], [], [], [] + last_key_frame_position = (keyframe_positions[-1]) + buffer - cn_frame_numbers = [] - cn_weights = [] - ipadapter_frame_numbers = [] - ipadapter_weights = [] - - last_key_frame_position = (keyframe_positions[-1]) + buffer - control_net = [] for i, (start, end) in enumerate(influence_ranges): + # set basic values batch_index_from, batch_index_to_excl = influence_ranges[i] ipadapter_strength_multiplier = relative_ipadapter_strength ipadapter_influence_multiplier = relative_ipadapter_influence + # Default values + revert_direction_at_midpoint = False + interpolation = "ease-in-out" + strength_from = strength_to = 1.0 + if i == 0: if buffer > 0: # First image with buffer image = images[0] - strength_from = strength_to = cn_strength_values[0][1] if len(cn_strength_values) > 0 else (1.0, 1.0) - revert_direction_at_midpoint = False - ipadapter_strength_multiplier = 1.0 + strength_from = strength_to = cn_strength_values[0][1] if len(cn_strength_values) > 0 else (1.0, 1.0) ipadapter_influence_multiplier = 1.0 interpolation = "ease-in-out" else: continue # Skip first image without buffer elif i == 1: # First image image = images[0] - strength_to, strength_from = cn_strength_values[0] if len(cn_strength_values) > 0 else (0.0, 1.0) - revert_direction_at_midpoint = False - interpolation = "ease-in" - - - + strength_to, strength_from = cn_strength_values[0] if len(cn_strength_values) > 0 else (0.0, 1.0) + interpolation = "ease-in" elif i == len(images): # Last image image = images[i-1] - strength_from, strength_to = cn_strength_values[i-1] if i-1 < len(cn_strength_values) else (0.0, 1.0) - revert_direction_at_midpoint = False - interpolation = "ease-out" - + strength_from, strength_to = cn_strength_values[i-1] if i-1 < len(cn_strength_values) else (0.0, 1.0) + interpolation = "ease-out" else: # Middle images image = images[i-1] - strength_from, strength_to = cn_strength_values[i-1] if i-1 < len(cn_strength_values) else (0.0, 1.0) + strength_from, strength_to = cn_strength_values[i-1] if i-1 < len(cn_strength_values) else (0.0, 1.0) revert_direction_at_midpoint = True - interpolation = "ease-in-out" - + # Import necessary modules latent_keyframe_interpolation_node = LatentKeyframeInterpolationNodeImport() - weights, frame_numbers, latent_keyframe, = latent_keyframe_interpolation_node.load_keyframe(batch_index_from,strength_from,batch_index_to_excl,strength_to,interpolation,revert_direction_at_midpoint,last_key_frame_position,i,len(influence_ranges),buffer) + scaled_soft_control_net_weights = ScaledSoftControlNetWeightsImport() + timestep_keyframe_node = TimestepKeyframeNodeImport() + control_net_loader = ControlNetLoaderAdvancedImport() + apply_advanced_control_net = AdvancedControlNetApplyImport() + ipadapter_application = IPAdapterApplyImport() + # Load keyframe and append frame numbers and weights + weights, frame_numbers, latent_keyframe = latent_keyframe_interpolation_node.load_keyframe( + batch_index_from, strength_from, batch_index_to_excl, strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position, i, len(influence_ranges), buffer) cn_frame_numbers.append(frame_numbers) cn_weights.append(weights) - scaled_soft_control_net_weights = ScaledSoftControlNetWeightsImport() - control_net_weights, _ = scaled_soft_control_net_weights.load_weights(soft_scaled_cn_weights_multiplier,False) + # Load weights and keyframe + control_net_weights, _ = scaled_soft_control_net_weights.load_weights(soft_scaled_cn_weights_multiplier, False) + timestep_keyframe = timestep_keyframe_node.load_keyframe(start_percent=0.0, control_net_weights=control_net_weights, t2i_adapter_weights=None, latent_keyframe=latent_keyframe, prev_timestep_keyframe=None)[0] - timestep_keyframe_node = TimestepKeyframeNodeImport() - timestep_keyframe, = timestep_keyframe_node.load_keyframe(start_percent=0.0,control_net_weights=control_net_weights,t2i_adapter_weights=None,latent_keyframe=latent_keyframe,prev_timestep_keyframe=None) + # Load and apply control net + control_net = control_net_loader.load_controlnet(control_net_name, timestep_keyframe)[0] + positive, negative = apply_advanced_control_net.apply_controlnet(positive, negative, control_net, image.unsqueeze(0), 1.0, 0.0, 1.0) - control_net_loader = ControlNetLoaderAdvancedImport() - control_net, = control_net_loader.load_controlnet(control_net_name, timestep_keyframe) + # Prepare image + prepped_image = prep_image(image=image.unsqueeze(0), interpolation="LANCZOS", crop_position="pad", sharpening=0.0)[0] - apply_advanced_control_net = AdvancedControlNetApplyImport() - positive, negative = apply_advanced_control_net.apply_controlnet(positive,negative,control_net,image.unsqueeze(0),1.0,0.0,1.0) - - prepped_image, = prep_image(image=image.unsqueeze(0), interpolation="LANCZOS", crop_position="pad", sharpening=0.0) - - ipadapter_application = IPAdapterApply() - - ipa_strength_from, ipa_strength_to = adjust_strength_values(strength_from, strength_to, ipadapter_strength_multiplier) - + # Adjust strength values and influence range + ipa_strength_from, ipa_strength_to = adjust_strength_values(strength_from, strength_to, ipadapter_strength_multiplier) ipa_batch_index_from, ipa_batch_index_to_excl = adjust_influence_range(batch_index_from, batch_index_to_excl, last_key_frame_position, ipadapter_influence_multiplier, buffer) - ipa_weights, ipa_frame_numbers = calculate_weights(ipa_batch_index_from,ipa_batch_index_to_excl,ipa_strength_from, ipa_strength_to,interpolation, revert_direction_at_midpoint, last_key_frame_position,i, len(influence_ranges),buffer) - - + # Calculate weights and append frame numbers and weights + ipa_weights, ipa_frame_numbers = calculate_weights(ipa_batch_index_from, ipa_batch_index_to_excl, ipa_strength_from, ipa_strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position, i, len(influence_ranges), buffer) ipadapter_frame_numbers.append(ipa_frame_numbers) ipadapter_weights.append(ipa_weights) - + # Create mask batch and apply ipadapter masks = create_mask_batch(last_key_frame_position, weights, frame_numbers) - - model, = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, clip_vision=clip_vision, image=prepped_image, weight_type="original", noise=ipadapter_noise, embeds=None, attn_mask=masks, start_at=0.0, end_at=1.0, unfold_batch=True) + model = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, clip_vision=clip_vision, image=prepped_image, weight_type="original", noise=ipadapter_noise, embeds=None, attn_mask=masks, start_at=0.0, end_at=1.0, unfold_batch=True)[0] comparison_diagram, = plot_weight_comparison(cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights, buffer) return comparison_diagram, positive, negative, model + # NODE MAPPING NODE_CLASS_MAPPINGS = { - # Combined "BatchCreativeInterpolation": BatchCreativeInterpolationNode - # "MaskGenerator": MaskGeneratorNode - # "FILMVFIImport": FILMVFINode - # Keyframes - # "TimestepKeyframe": TimestepKeyframeNodeImport, - # "LatentKeyframeImport": LatentKeyframeNodeImport, - # "LatentKeyframeGroupImport": LatentKeyframeGroupImportNode, - # "LatentKeyframeBatchedGroupImport": LatentKeyframeBatchedGroupNodeImport, - # "LatentKeyframeTiming": LatentKeyframeInterpolationNodeImport, - # Loaders - # "ControlNetLoaderAdvancedImport": ControlNetLoaderAdvancedImport, - # "DiffControlNetLoaderAdvancedImport": DiffControlNetLoaderAdvancedImport, - # Conditioning - # "ACN_AdvancedControlNetApplyImport": AdvancedControlNetApplyImport, - # Weights - # "ScaledSoftControlNetWeightsImport": ScaledSoftControlNetWeightsImport, - # "SoftControlNetWeights": SoftControlNetWeights, - # "CustomControlNetWeights": CustomControlNetWeights, - # "SoftT2IAdapterWeights": SoftT2IAdapterWeights, - # "CustomT2IAdapterWeights": CustomT2IAdapterWeights, - # Image - # "LoadImagesFromDirectory": LoadImagesFromDirectory } -NODE_DISPLAY_NAME_MAPPINGS = { - # Combined - "BatchCreativeInterpolation": "Batch Creative Interpolation 🎞️🅟🅞🅜" - # "MaskGenerator": "Mask Generator 🎞️🅟🅞🅜" - # Keyframes - # "TimestepKeyframe": "Timestep Keyframe 🎞️🅟🅞🅜", - # "LatentKeyframe": "Latent Keyframe 🛂🅐🅒🅝", - # "LatentKeyframeGroupImport": "Latent Keyframe Group 🛂🅐🅒🅝", - # "LatentKeyframeBatchedGroup": "Latent Keyframe Batched Group 🛂🅐🅒🅝", - # "LatentKeyframeTiming": "Latent Keyframe Interpolation 🛂🅐🅒🅝", - # Loaders - # "ControlNetLoaderAdvancedImport": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", - # "DiffControlNetLoaderAdvancedImport": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", - # Conditioning - # "ACN_AdvancedControlNetApplyImport": "Apply Advanced ControlNet 🛂🅐🅒🅝", - # Weights - # "ScaledSoftControlNetWeightsImport": "Scaled Soft ControlNet Weights 🛂🅐🅒🅝", - # "SoftControlNetWeights": "Soft ControlNet Weights 🛂🅐🅒🅝", - # "CustomControlNetWeights": "Custom ControlNet Weights 🛂🅐🅒🅝", - # "SoftT2IAdapterWeights": "Soft T2IAdapter Weights 🛂🅐🅒🅝", - # "CustomT2IAdapterWeights": "Custom T2IAdapter Weights 🛂🅐🅒🅝", - # Image - # "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" +NODE_DISPLAY_NAME_MAPPINGS = { + "BatchCreativeInterpolation": "Batch Creative Interpolation 🎞️🅢🅜" } diff --git a/__init__.py b/__init__.py index e70bf90..52ce1cd 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,3 @@ -from .control.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .SteerableMotion import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/control/.DS_Store b/control/.DS_Store deleted file mode 100644 index 5008ddfcf53c02e82d7eee2e57c38e5672ef89f6..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 6148 zcmeH~Jr2S!425mzP>H1@V-^m;4Wg<&0T*E43hX&L&p$$qDprKhvt+--jT7}7np#A3 zem<@ulZcFPQ@L2!n>{z**++&mCkOWA81W14cNZlEfg7;MkzE(HCqgga^y>{tEnwC%0;vJ&^%eQ zLs35+`xjp>T0 0: - limit_images = True - image_count = 0 - - for image_path in dir_files: - if os.path.isdir(image_path): - continue - if limit_images and image_count >= image_load_cap: - break - i = Image.open(image_path) - i = ImageOps.exif_transpose(i) - image = i.convert("RGB") - image = np.array(image).astype(np.float32) / 255.0 - image = torch.from_numpy(image)[None,] - if 'A' in i.getbands(): - mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 - mask = 1. - torch.from_numpy(mask) - else: - mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") - images.append(image) - masks.append(mask) - image_count += 1 - - if len(images) == 0: - raise FileNotFoundError(f"No images could be loaded from directory '{directory}'.") - - return (torch.cat(images, dim=0), torch.stack(masks, dim=0), image_count) diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py deleted file mode 100644 index faf932d..0000000 --- a/control/latent_keyframe_nodes.py +++ /dev/null @@ -1,296 +0,0 @@ -from typing import Union -import numpy as np -from collections.abc import Iterable - -from .control import LatentKeyframeImport, LatentKeyframeGroupImport -from .logger import logger - - -class LatentKeyframeNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.00001}, ), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - } - } - - RETURN_TYPES = ("LATENT_KEYFRAME", ) - FUNCTION = "load_keyframe" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, - batch_index: int, - strength: float, - prev_latent_keyframe: LatentKeyframeGroupImport=None): - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - keyframe = LatentKeyframeImport(batch_index, strength) - prev_latent_keyframe.add(keyframe) - return (prev_latent_keyframe,) - - -class LatentKeyframeGroupNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "index_strengths": ("STRING", {"multiline": True, "default": ""}), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - "latent_optional": ("LATENT", ), - } - } - - RETURN_TYPES = ("LATENT_KEYFRAME", ) - FUNCTION = "load_keyframes" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def validate_index(self, index: int, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: - # if part of range, do nothing - if is_range: - return index - # otherwise, validate index - # validate not out of range - only when latent_count is passed in - if latent_count > 0 and index > latent_count-1: - raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.") - # if negative, validate not out of range - if index < 0: - if not allow_negative: - raise IndexError(f"Negative indeces not allowed, but was {index}.") - conv_index = latent_count+index - if conv_index < 0: - raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.") - index = conv_index - return index - - def convert_to_index_int(self, raw_index: str, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: - try: - return self.validate_index(int(raw_index), latent_count=latent_count, is_range=is_range, allow_negative=allow_negative) - except ValueError as e: - raise ValueError(f"index '{raw_index}' must be an integer.", e) - - def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframeImport]: - if not latent_indeces: - return set() - all_indeces = [i for i in range(0, latent_count)] - allow_negative = latent_count > 0 - chosen_indeces = set() - # parse string - allow positive ints, negative ints, and ranges separated by ':' - groups = latent_indeces.split(",") - groups = [g.strip() for g in groups] - for g in groups: - # parse strengths - default to 1.0 if no strength given - strength = 1.0 - if '=' in g: - g, strength_str = g.split("=", 1) - g = g.strip() - try: - strength = float(strength_str.strip()) - except ValueError as e: - raise ValueError(f"strength '{strength_str}' must be a float.", e) - if strength < 0: - raise ValueError(f"Strength '{strength}' cannot be negative.") - # parse range of indeces (e.g. 2:16) - if ':' in g: - index_range = g.split(":", 1) - index_range = [r.strip() for r in index_range] - start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative) - end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative) - for i in all_indeces[start_index:end_index]: - chosen_indeces.add(LatentKeyframImport(i, strength)) - # parse individual indeces - else: - chosen_indeces.add(LatentKeyframeImport(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength)) - return chosen_indeces - - def load_keyframes(self, - index_strengths: str, - prev_latent_keyframe: LatentKeyframeGroupImport=None, - latent_image_opt=None): - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - curr_latent_keyframe = LatentKeyframeGroupImport() - - latent_count = -1 - if latent_image_opt: - latent_count = latent_image_opt['samples'].size()[0] - latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count) - - for latent_keyframe in latent_keyframes: - logger.info(f"keyframe {latent_keyframe.batch_index}:{latent_keyframe.strength}") - curr_latent_keyframe.add(latent_keyframe) - - for latent_keyframe in prev_latent_keyframe.keyframes: - curr_latent_keyframe.add(latent_keyframe) - - return (curr_latent_keyframe,) - - -class LatentKeyframeInterpolationNodeImport: - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), - "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), - "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"], ), - "revert_direction_at_midpoint": ("BOOLEAN", {"default": False}), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - } - } - - RETURN_TYPES = ("LATENT_KEYFRAME", ) - FUNCTION = "load_keyframe" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, - batch_index_from: int, - strength_from: float, - batch_index_to_excl: int, - strength_to: float, - interpolation: str, - revert_direction_at_midpoint: bool=False, - last_key_frame_position: int=0, - i=0, - number_of_items=0, - buffer=0, - prev_latent_keyframe: LatentKeyframeGroupImport=None): - - - - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - - curr_latent_keyframe = LatentKeyframeGroupImport() - - weights, frame_numbers = calculate_weights(batch_index_from, batch_index_to_excl, strength_from, strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position,i,number_of_items, buffer) - - for i, frame_number in enumerate(frame_numbers): - keyframe = LatentKeyframeImport(frame_number, float(weights[i])) - logger.info(f"keyframe {frame_number}:{weights[i]}") - curr_latent_keyframe.add(keyframe) - - for latent_keyframe in prev_latent_keyframe.keyframes: - curr_latent_keyframe.add(latent_keyframe) - - - return (weights, frame_numbers, curr_latent_keyframe,) - - -def calculate_weights(batch_index_from, batch_index_to, strength_from, strength_to, interpolation,revert_direction_at_midpoint, last_key_frame_position,i, number_of_items,buffer): - print("<-------- LOOK HERE s") - print("batch_index_from",batch_index_from) - print("batch_index_to",batch_index_to) - print("strength_from",strength_from) - print("strength_to",strength_to) - print("interpolation",interpolation) - print("revert_direction_at_midpoint",revert_direction_at_midpoint) - print("last_key_frame_position",last_key_frame_position) - print("i",i) - print("number_of_items",number_of_items) - # Initialize variables based on the position of the keyframe - range_start = batch_index_from - range_end = batch_index_to - # if it's the first value, set influence range from 1.0 to 0.0 - if buffer > 0: - if i == 0: - range_start = 0 - elif i == 1: - range_start = buffer - else: - if i == 1: - range_start = 0 - - if i == number_of_items - 1: - range_end = last_key_frame_position - - steps = range_end - range_start - diff = strength_to - strength_from - - # Calculate index for interpolation - index = np.linspace(0, 1, steps // 2 + 1) if revert_direction_at_midpoint else np.linspace(0, 1, steps) - - # Calculate weights based on interpolation type - if interpolation == "linear": - weights = np.linspace(strength_from, strength_to, len(index)) - elif interpolation == "ease-in": - weights = diff * np.power(index, 2) + strength_from - elif interpolation == "ease-out": - weights = diff * (1 - np.power(1 - index, 2)) + strength_from - elif interpolation == "ease-in-out": - weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from - - # If it's a middle keyframe, mirror the weights - if revert_direction_at_midpoint: - weights = np.concatenate([weights, weights[::-1]]) - - # Generate frame numbers - frame_numbers = np.arange(range_start, range_start + len(weights)) - - # "Dropper" component: For keyframes with negative start, drop the weights - if range_start < 0 and i > 0: - drop_count = abs(range_start) - weights = weights[drop_count:] - frame_numbers = frame_numbers[drop_count:] - - # Dropper component: for keyframes a range_End is greater than last_key_frame_position, drop the weights - if range_end > last_key_frame_position and i < number_of_items - 1: - drop_count = range_end - last_key_frame_position - weights = weights[:-drop_count] - frame_numbers = frame_numbers[:-drop_count] - - return weights, frame_numbers - - - -class LatentKeyframeBatchedGroupNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001}), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - } - } - - RETURN_TYPES = ("LATENT_KEYFRAME", ) - FUNCTION = "load_keyframe" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroupImport=None): - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - curr_latent_keyframe = LatentKeyframeGroupImport() - - # if received a normal float input, do nothing - if type(strengths) in (float, int): - logger.info("No batched strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.") - # if iterable, attempt to create LatentKeyframes with chosen strengths - elif isinstance(strengths, Iterable): - for idx, strength in enumerate(strengths): - keyframe = LatentKeyframeImport(idx, strength) - curr_latent_keyframe.add(keyframe) - logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}") - else: - raise ValueError(f"Expected strengths to be an iterable input, but was {type(strengths).__repr__}.") - - # replace values with prev_latent_keyframes - for latent_keyframe in prev_latent_keyframe.keyframes: - curr_latent_keyframe.add(latent_keyframe) - - return (curr_latent_keyframe,) diff --git a/control/logger.py b/control/logger.py deleted file mode 100644 index b23b82f..0000000 --- a/control/logger.py +++ /dev/null @@ -1,36 +0,0 @@ -import sys -import copy -import logging - - -class ColoredFormatter(logging.Formatter): - COLORS = { - "DEBUG": "\033[0;36m", # CYAN - "INFO": "\033[0;32m", # GREEN - "WARNING": "\033[0;33m", # YELLOW - "ERROR": "\033[0;31m", # RED - "CRITICAL": "\033[0;37;41m", # WHITE ON RED - "RESET": "\033[0m", # RESET COLOR - } - - def format(self, record): - colored_record = copy.copy(record) - levelname = colored_record.levelname - seq = self.COLORS.get(levelname, self.COLORS["RESET"]) - colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" - return super().format(colored_record) - - -# Create a new logger -logger = logging.getLogger("Advanced-ControlNet") -logger.propagate = False - -# Add handler if we don't have one. -if not logger.handlers: - handler = logging.StreamHandler(sys.stdout) - handler.setFormatter(ColoredFormatter("[%(name)s] - %(levelname)s - %(message)s")) - logger.addHandler(handler) - -# Configure logger -loglevel = logging.INFO -logger.setLevel(loglevel) diff --git a/control/resampler.py b/control/resampler.py deleted file mode 100644 index 4521c8c..0000000 --- a/control/resampler.py +++ /dev/null @@ -1,121 +0,0 @@ -# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py -import math - -import torch -import torch.nn as nn - - -# FFN -def FeedForward(dim, mult=4): - inner_dim = int(dim * mult) - return nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, inner_dim, bias=False), - nn.GELU(), - nn.Linear(inner_dim, dim, bias=False), - ) - - -def reshape_tensor(x, heads): - bs, length, width = x.shape - #(bs, length, width) --> (bs, length, n_heads, dim_per_head) - x = x.view(bs, length, heads, -1) - # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) - x = x.transpose(1, 2) - # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) - x = x.reshape(bs, heads, length, -1) - return x - - -class PerceiverAttention(nn.Module): - def __init__(self, *, dim, dim_head=64, heads=8): - super().__init__() - self.scale = dim_head**-0.5 - self.dim_head = dim_head - self.heads = heads - inner_dim = dim_head * heads - - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) - - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) - self.to_out = nn.Linear(inner_dim, dim, bias=False) - - - def forward(self, x, latents): - """ - Args: - x (torch.Tensor): image features - shape (b, n1, D) - latent (torch.Tensor): latent features - shape (b, n2, D) - """ - x = self.norm1(x) - latents = self.norm2(latents) - - b, l, _ = latents.shape - - q = self.to_q(latents) - kv_input = torch.cat((x, latents), dim=-2) - k, v = self.to_kv(kv_input).chunk(2, dim=-1) - - q = reshape_tensor(q, self.heads) - k = reshape_tensor(k, self.heads) - v = reshape_tensor(v, self.heads) - - # attention - scale = 1 / math.sqrt(math.sqrt(self.dim_head)) - weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - out = weight @ v - - out = out.permute(0, 2, 1, 3).reshape(b, l, -1) - - return self.to_out(out) - - -class Resampler(nn.Module): - def __init__( - self, - dim=1024, - depth=8, - dim_head=64, - heads=16, - num_queries=8, - embedding_dim=768, - output_dim=1024, - ff_mult=4, - ): - super().__init__() - - self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) - - self.proj_in = nn.Linear(embedding_dim, dim) - - self.proj_out = nn.Linear(dim, output_dim) - self.norm_out = nn.LayerNorm(output_dim) - - self.layers = nn.ModuleList([]) - for _ in range(depth): - self.layers.append( - nn.ModuleList( - [ - PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), - FeedForward(dim=dim, mult=ff_mult), - ] - ) - ) - - def forward(self, x): - - latents = self.latents.repeat(x.size(0), 1, 1) - - x = self.proj_in(x) - - for attn, ff in self.layers: - latents = attn(x, latents) + latents - latents = ff(latents) + latents - - latents = self.proj_out(latents) - return self.norm_out(latents) \ No newline at end of file diff --git a/control/weight_nodes.py b/control/weight_nodes.py deleted file mode 100644 index b84d6be..0000000 --- a/control/weight_nodes.py +++ /dev/null @@ -1,157 +0,0 @@ -from .control import TimestepKeyframeImport, TimestepKeyframeGroupImport -from .logger import logger - - -def get_properly_arranged_t2i_weights(initial_weights: list[float]): - new_weights = [] - new_weights.extend([initial_weights[0]]*3) - new_weights.extend([initial_weights[1]]*3) - new_weights.extend([initial_weights[2]]*3) - new_weights.extend([initial_weights[3]]*3) - return new_weights - - -class ScaledSoftControlNetWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - - def load_weights(self, base_multiplier, flip_weights): - weights = [(base_multiplier ** float(12 - i)) for i in range(13)] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_net_weights=weights))) - - -class SoftControlNetWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "weight_00": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_01": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_02": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_03": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_04": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_05": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_06": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_07": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_08": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_09": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - - def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): - weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_net_weights=weights))) - - -class CustomControlNetWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_04": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_05": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_06": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_07": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_08": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_09": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - } - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - - def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): - weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_net_weights=weights))) - - -class SoftT2IAdapterWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - - def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): - weights = [weight_00, weight_01, weight_02, weight_03] - if flip_weights: - weights.reverse() - weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(t2i_adapter_weights=weights))) - - -class CustomT2IAdapterWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - - def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): - weights = [weight_00, weight_01, weight_02, weight_03] - if flip_weights: - weights.reverse() - weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(t2i_adapter_weights=weights))) diff --git a/control/control.py b/imports/AdvancedControlNet.py similarity index 51% rename from control/control.py rename to imports/AdvancedControlNet.py index e559bf8..8e82867 100644 --- a/control/control.py +++ b/imports/AdvancedControlNet.py @@ -1,23 +1,23 @@ from typing import Union -from torch import Tensor -import torch +from collections.abc import Iterable +import folder_paths +import torch +import numpy as np +from torch import Tensor +from comfy.controlnet import ControlNet, T2IAdapter,broadcast_image_to import comfy.utils import comfy.controlnet as comfy_cn -from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to - ControlNetWeightsTypeImport = list[float] T2IAdapterWeightsTypeImport = list[float] - + class LatentKeyframeImport: def __init__(self, batch_index: int, strength: float) -> None: self.batch_index = batch_index self.strength = strength - -# always maintain sorted state (by batch_index of LatentKeyframe) class LatentKeyframeGroupImport: def __init__(self) -> None: self.keyframes: list[LatentKeyframeImport] = [] @@ -46,7 +46,6 @@ class LatentKeyframeGroupImport: def is_empty(self) -> bool: return len(self.keyframes) == 0 - class TimestepKeyframeImport: def __init__(self, start_percent: float = 0.0, @@ -64,9 +63,7 @@ class TimestepKeyframeImport: @classmethod def default(cls) -> 'TimestepKeyframeImport': return cls(0.0) - - -# always maintain sorted state (by start_percent of TimestepKeyFrame) + class TimestepKeyframeGroupImport: def __init__(self) -> None: self.keyframes: list[TimestepKeyframeImport] = [] @@ -103,56 +100,220 @@ class TimestepKeyframeGroupImport: return group -# used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function -def control_merge_inject(self, control_input, control_output, control_prev, output_dtype): - out = {'input':[], 'middle':[], 'output': []} - if control_input is not None: - for i in range(len(control_input)): - key = 'input' - x = control_input[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - out[key].insert(0, x) +class AdvancedControlNetApplyImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "control_net": ("CONTROL_NET", ), + "image": ("IMAGE", ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) + }, + "optional": { + "mask_optional": ("MASK", ), + } + } - if control_output is not None: - for i in range(len(control_output)): - if i == (len(control_output) - 1): - key = 'middle' - index = 0 + RETURN_TYPES = ("CONDITIONING","CONDITIONING") + RETURN_NAMES = ("positive", "negative") + FUNCTION = "apply_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" + + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=None): + if strength == 0: + return (positive, negative) + + control_hint = image.movedim(-1,1) + cnets = {} + + out = [] + for conditioning in [positive, negative]: + c = [] + + for t in conditioning: + d = t[1].copy() + + prev_cnet = d.get('control', None) + if prev_cnet in cnets: + c_net = cnets[prev_cnet] + + else: + c_net = control_net.copy().set_cond_hint(control_hint, strength, (start_percent, end_percent)) + # set cond hint mask + if mask_optional is not None: + if is_advanced_controlnet(c_net): + # if not in the form of a batch, make it so + if len(mask_optional.shape) < 3: + mask_optional = mask_optional.unsqueeze(0) + c_net.set_cond_hint_mask(mask_optional) + c_net.set_previous_controlnet(prev_cnet) + cnets[prev_cnet] = c_net + + d['control'] = c_net + d['control_apply_to_uncond'] = False + n = [t[0], d] + c.append(n) + out.append(c) + return (out[0], out[1]) + +class LatentKeyframeGroupNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "index_strengths": ("STRING", {"multiline": True, "default": ""}), + }, + "optional": { + "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "latent_optional": ("LATENT", ), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframes" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def validate_index(self, index: int, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: + # if part of range, do nothing + if is_range: + return index + # otherwise, validate index + # validate not out of range - only when latent_count is passed in + if latent_count > 0 and index > latent_count-1: + raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.") + # if negative, validate not out of range + if index < 0: + if not allow_negative: + raise IndexError(f"Negative indeces not allowed, but was {index}.") + conv_index = latent_count+index + if conv_index < 0: + raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.") + index = conv_index + return index + + def convert_to_index_int(self, raw_index: str, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: + try: + return self.validate_index(int(raw_index), latent_count=latent_count, is_range=is_range, allow_negative=allow_negative) + except ValueError as e: + raise ValueError(f"index '{raw_index}' must be an integer.", e) + + def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframeImport]: + if not latent_indeces: + return set() + all_indeces = [i for i in range(0, latent_count)] + allow_negative = latent_count > 0 + chosen_indeces = set() + # parse string - allow positive ints, negative ints, and ranges separated by ':' + groups = latent_indeces.split(",") + groups = [g.strip() for g in groups] + for g in groups: + # parse strengths - default to 1.0 if no strength given + strength = 1.0 + if '=' in g: + g, strength_str = g.split("=", 1) + g = g.strip() + try: + strength = float(strength_str.strip()) + except ValueError as e: + raise ValueError(f"strength '{strength_str}' must be a float.", e) + if strength < 0: + raise ValueError(f"Strength '{strength}' cannot be negative.") + # parse range of indeces (e.g. 2:16) + if ':' in g: + index_range = g.split(":", 1) + index_range = [r.strip() for r in index_range] + start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative) + end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative) + for i in all_indeces[start_index:end_index]: + chosen_indeces.add(LatentKeyframeImport(i, strength)) + # parse individual indeces else: - key = 'output' - index = i - x = control_output[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) + chosen_indeces.add(LatentKeyframeImport(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength)) + return chosen_indeces - if self.global_average_pooling: - x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) + def load_keyframes(self, + index_strengths: str, + prev_latent_keyframe: LatentKeyframeGroupImport=None, + latent_image_opt=None): + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + curr_latent_keyframe = LatentKeyframeGroupImport() - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) + latent_count = -1 + if latent_image_opt: + latent_count = latent_image_opt['samples'].size()[0] + latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count) - out[key].append(x) - if control_prev is not None: - for x in ['input', 'middle', 'output']: - o = out[x] - for i in range(len(control_prev[x])): - prev_val = control_prev[x][i] - if i >= len(o): - o.append(prev_val) - elif prev_val is not None: - if o[i] is None: - o[i] = prev_val - else: - o[i] += prev_val - return out + for latent_keyframe in latent_keyframes: + + curr_latent_keyframe.add(latent_keyframe) + + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + return (curr_latent_keyframe,) + +class LatentKeyframeInterpolationNodeImport: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), + "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), + "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), + "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), + "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"], ), + "revert_direction_at_midpoint": ("BOOLEAN", {"default": False}), + }, + "optional": { + "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframe" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + batch_index_from: int, + strength_from: float, + batch_index_to_excl: int, + strength_to: float, + interpolation: str, + revert_direction_at_midpoint: bool=False, + last_key_frame_position: int=0, + i=0, + number_of_items=0, + buffer=0, + prev_latent_keyframe: LatentKeyframeGroupImport=None): + + + + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + + curr_latent_keyframe = LatentKeyframeGroupImport() + + weights, frame_numbers = calculate_weights(batch_index_from, batch_index_to_excl, strength_from, strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position,i,number_of_items, buffer) + + for i, frame_number in enumerate(frame_numbers): + keyframe = LatentKeyframeImport(frame_number, float(weights[i])) + curr_latent_keyframe.add(keyframe) + + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + + return (weights, frame_numbers, curr_latent_keyframe,) class ControlNetAdvancedImport(ControlNet): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroupImport, global_average_pooling=False, device=None): @@ -298,6 +459,118 @@ class ControlNetAdvancedImport(ControlNet): self.full_latent_length = 0 self.context_length = 0 +class ControlNetLoaderAdvancedImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + }, + "optional": { + "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + + def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroupImport=None): + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + controlnet = load_controlnet(controlnet_path, timestep_keyframe) + return (controlnet,) + +class TimestepKeyframeNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + }, + "optional": { + "control_net_weights": ("CONTROL_NET_WEIGHTS", ), + "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), + "latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) + FUNCTION = "load_keyframe" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + start_percent: float, + control_net_weights: ControlNetWeightsTypeImport=None, + t2i_adapter_weights: T2IAdapterWeightsTypeImport=None, + latent_keyframe: LatentKeyframeGroupImport=None, + prev_timestep_keyframe: TimestepKeyframeGroupImport=None): + if not prev_timestep_keyframe: + prev_timestep_keyframe = TimestepKeyframeGroupImport() + keyframe = TimestepKeyframeImport(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe) + prev_timestep_keyframe.add(keyframe) + return (prev_timestep_keyframe,) + +class ScaledSoftControlNetWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self, base_multiplier, flip_weights): + weights = [(base_multiplier ** float(12 - i)) for i in range(13)] + if flip_weights: + weights.reverse() + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_net_weights=weights))) + +class LatentKeyframeBatchedGroupNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001}), + }, + "optional": { + "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframe" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroupImport=None): + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + curr_latent_keyframe = LatentKeyframeGroupImport() + + # if received a normal float input, do nothing + if type(strengths) in (float, int): + print("No batched strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.") + # if iterable, attempt to create LatentKeyframes with chosen strengths + elif isinstance(strengths, Iterable): + for idx, strength in enumerate(strengths): + keyframe = LatentKeyframeImport(idx, strength) + curr_latent_keyframe.add(keyframe) + else: + raise ValueError(f"Expected strengths to be an iterable input, but was {type(strengths).__repr__}.") + + # replace values with prev_latent_keyframes + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + return (curr_latent_keyframe,) class T2IAdapterAdvancedImport(T2IAdapter): def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroupImport, channels_in, device=None): @@ -351,6 +624,65 @@ class T2IAdapterAdvancedImport(T2IAdapter): self.context_length = 0 +def is_advanced_controlnet(input_object): + return isinstance(input_object, ControlNetAdvancedImport) or isinstance(input_object, T2IAdapterAdvancedImport) + +def control_merge_inject(self, control_input, control_output, control_prev, output_dtype): + out = {'input':[], 'middle':[], 'output': []} + + if control_input is not None: + for i in range(len(control_input)): + key = 'input' + x = control_input[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) + + x *= self.strength * self.weights[i] + if x.dtype != output_dtype: + x = x.to(output_dtype) + out[key].insert(0, x) + + if control_output is not None: + for i in range(len(control_output)): + if i == (len(control_output) - 1): + key = 'middle' + index = 0 + else: + key = 'output' + index = i + x = control_output[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) + + if self.global_average_pooling: + x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) + + x *= self.strength * self.weights[i] + if x.dtype != output_dtype: + x = x.to(output_dtype) + + out[key].append(x) + if control_prev is not None: + for x in ['input', 'middle', 'output']: + o = out[x] + for i in range(len(control_prev[x])): + prev_val = control_prev[x][i] + if i >= len(o): + o.append(prev_val) + elif prev_val is not None: + if o[i] is None: + o[i] = prev_val + else: + o[i] += prev_val + return out + +def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): + mask = mask.clone() + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") + if match_dim1: + mask = torch.cat([mask] * shape[1], dim=1) + return mask + def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroupImport=None, model=None): control = comfy_cn.load_controlnet(ckpt_path, model=model) # if exactly ControlNet returned, transform it into ControlNetAdvanced @@ -363,15 +695,57 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroupImport=No # TODO add ControlLoraAdvanced return control +def calculate_weights(batch_index_from, batch_index_to, strength_from, strength_to, interpolation,revert_direction_at_midpoint, last_key_frame_position,i, number_of_items,buffer): -def is_advanced_controlnet(input_object): - return isinstance(input_object, ControlNetAdvancedImport) or isinstance(input_object, T2IAdapterAdvancedImport) + # Initialize variables based on the position of the keyframe + range_start = batch_index_from + range_end = batch_index_to + # if it's the first value, set influence range from 1.0 to 0.0 + if buffer > 0: + if i == 0: + range_start = 0 + elif i == 1: + range_start = buffer + else: + if i == 1: + range_start = 0 + + if i == number_of_items - 1: + range_end = last_key_frame_position + steps = range_end - range_start + diff = strength_to - strength_from -# adapted from comfy/sample.py -def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): - mask = mask.clone() - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") - if match_dim1: - mask = torch.cat([mask] * shape[1], dim=1) - return mask + # Calculate index for interpolation + index = np.linspace(0, 1, steps // 2 + 1) if revert_direction_at_midpoint else np.linspace(0, 1, steps) + + # Calculate weights based on interpolation type + if interpolation == "linear": + weights = np.linspace(strength_from, strength_to, len(index)) + elif interpolation == "ease-in": + weights = diff * np.power(index, 2) + strength_from + elif interpolation == "ease-out": + weights = diff * (1 - np.power(1 - index, 2)) + strength_from + elif interpolation == "ease-in-out": + weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from + + # If it's a middle keyframe, mirror the weights + if revert_direction_at_midpoint: + weights = np.concatenate([weights, weights[::-1]]) + + # Generate frame numbers + frame_numbers = np.arange(range_start, range_start + len(weights)) + + # "Dropper" component: For keyframes with negative start, drop the weights + if range_start < 0 and i > 0: + drop_count = abs(range_start) + weights = weights[drop_count:] + frame_numbers = frame_numbers[drop_count:] + + # Dropper component: for keyframes a range_End is greater than last_key_frame_position, drop the weights + if range_end > last_key_frame_position and i < number_of_items - 1: + drop_count = range_end - last_key_frame_position + weights = weights[:-drop_count] + frame_numbers = frame_numbers[:-drop_count] + + return weights, frame_numbers \ No newline at end of file diff --git a/control/IPAdapterPlus.py b/imports/IPAdapterPlus.py similarity index 77% rename from control/IPAdapterPlus.py rename to imports/IPAdapterPlus.py index a2300bc..58f4eb9 100644 --- a/control/IPAdapterPlus.py +++ b/imports/IPAdapterPlus.py @@ -14,8 +14,6 @@ from PIL import Image import torch.nn.functional as F import torchvision.transforms as TT -from .resampler import Resampler - # set the models directory backward compatible GLOBAL_MODELS_DIR = os.path.join(folder_paths.models_dir, "ipadapter") MODELS_DIR = GLOBAL_MODELS_DIR if os.path.isdir(GLOBAL_MODELS_DIR) else os.path.join(os.path.dirname(os.path.realpath(__file__)), "models") @@ -24,7 +22,7 @@ if "ipadapter" not in folder_paths.folder_names_and_paths: else: folder_paths.folder_names_and_paths["ipadapter"][1].update(folder_paths.supported_pt_extensions) -class MLPProjModel(torch.nn.Module): +class MLPProjModelImport(torch.nn.Module): """SD model with image prompt""" def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024): super().__init__() @@ -40,7 +38,7 @@ class MLPProjModel(torch.nn.Module): clip_extra_context_tokens = self.proj(image_embeds) return clip_extra_context_tokens -class ImageProjModel(nn.Module): +class ImageProjModelImport(nn.Module): def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4): super().__init__() @@ -55,7 +53,7 @@ class ImageProjModel(nn.Module): clip_extra_context_tokens = self.norm(clip_extra_context_tokens) return clip_extra_context_tokens -class To_KV(nn.Module): +class To_KVImport(nn.Module): def __init__(self, state_dict): super().__init__() @@ -64,6 +62,73 @@ class To_KV(nn.Module): self.to_kvs[key.replace(".weight", "").replace(".", "_")] = nn.Linear(value.shape[1], value.shape[0], bias=False) self.to_kvs[key.replace(".weight", "").replace(".", "_")].weight.data = value +def FeedForward(dim, mult=4): + inner_dim = int(dim * mult) + return nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, inner_dim, bias=False), + nn.GELU(), + nn.Linear(inner_dim, dim, bias=False), + ) + + +class PerceiverAttention(nn.Module): + def __init__(self, *, dim, dim_head=64, heads=8): + super().__init__() + self.scale = dim_head**-0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, l, _ = latents.shape + + q = self.to_q(latents) + kv_input = torch.cat((x, latents), dim=-2) + k, v = self.to_kv(kv_input).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, l, -1) + + return self.to_out(out) + +def reshape_tensor(x, heads): + bs, length, width = x.shape + #(bs, length, width) --> (bs, length, n_heads, dim_per_head) + x = x.view(bs, length, heads, -1) + # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) + x = x.transpose(1, 2) + # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) + x = x.reshape(bs, heads, length, -1) + return x + def set_model_patch_replace(model, patch_kwargs, key): to = model.model_options["transformer_options"] if "patches_replace" not in to: @@ -71,7 +136,7 @@ def set_model_patch_replace(model, patch_kwargs, key): if "attn2" not in to["patches_replace"]: to["patches_replace"]["attn2"] = {} if key not in to["patches_replace"]["attn2"]: - patch = CrossAttentionPatch(**patch_kwargs) + patch = CrossAttentionPatchImport(**patch_kwargs) to["patches_replace"]["attn2"][key] = patch else: to["patches_replace"]["attn2"][key].set_new_condition(**patch_kwargs) @@ -161,7 +226,7 @@ def contrast_adaptive_sharpening(image, amount): return (output) -class IPAdapter(nn.Module): +class IPAdapterImport(nn.Module): def __init__(self, ipadapter_model, cross_attention_dim=1024, output_cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4, is_sdxl=False, is_plus=False, is_full=False): super().__init__() @@ -174,10 +239,10 @@ class IPAdapter(nn.Module): self.image_proj_model = self.init_proj() if not is_plus else self.init_proj_plus() self.image_proj_model.load_state_dict(ipadapter_model["image_proj"]) - self.ip_layers = To_KV(ipadapter_model["ip_adapter"]) + self.ip_layers = To_KVImport(ipadapter_model["ip_adapter"]) def init_proj(self): - image_proj_model = ImageProjModel( + image_proj_model = ImageProjModelImport( cross_attention_dim=self.cross_attention_dim, clip_embeddings_dim=self.clip_embeddings_dim, clip_extra_context_tokens=self.clip_extra_context_tokens @@ -186,12 +251,12 @@ class IPAdapter(nn.Module): def init_proj_plus(self): if self.is_full: - image_proj_model = MLPProjModel( + image_proj_model = MLPProjModelImport( cross_attention_dim=self.cross_attention_dim, clip_embeddings_dim=self.clip_embeddings_dim ) else: - image_proj_model = Resampler( + image_proj_model = ResamplerImport( dim=self.cross_attention_dim, depth=4, dim_head=64, @@ -209,7 +274,7 @@ class IPAdapter(nn.Module): uncond_image_prompt_embeds = self.image_proj_model(clip_embed_zeroed) return image_prompt_embeds, uncond_image_prompt_embeds -class CrossAttentionPatch: +class CrossAttentionPatchImport: # forward for patching def __init__(self, weight, ipadapter, device, dtype, number, cond, uncond, weight_type, mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False): self.weights = [weight] @@ -368,7 +433,7 @@ class CrossAttentionPatch: return out.to(dtype=org_dtype) -class IPAdapterModelLoader: +class IPAdapterModelLoaderImport: @classmethod def INPUT_TYPES(s): return {"required": { "ipadapter_file": (folder_paths.get_filename_list("ipadapter"), )}} @@ -397,7 +462,7 @@ class IPAdapterModelLoader: return (model,) -class IPAdapterApply: +class IPAdapterApplyImport: @classmethod def INPUT_TYPES(s): return { @@ -463,7 +528,7 @@ class IPAdapterApply: clip_embeddings_dim = clip_embed.shape[-1] - self.ipadapter = IPAdapter( + self.ipadapter = IPAdapterImport( ipadapter, cross_attention_dim=cross_attention_dim, output_cross_attention_dim=output_cross_attention_dim, @@ -528,7 +593,6 @@ class IPAdapterApply: return (work_model, ) - def prep_image(image, interpolation="LANCZOS", crop_position="center", sharpening=0.0): _, oh, ow, _ = image.shape output = image.permute([0,3,1,2]) @@ -574,178 +638,47 @@ def prep_image(image, interpolation="LANCZOS", crop_position="center", sharpenin return (output,) -class IPAdapterEncoder: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "clip_vision": ("CLIP_VISION",), - "image_1": ("IMAGE",), - "ipadapter_plus": ("BOOLEAN", { "default": False }), - "noise": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01 }), - "weight_1": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), - }, - "optional": { - "image_2": ("IMAGE",), - "image_3": ("IMAGE",), - "image_4": ("IMAGE",), - "weight_2": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), - "weight_3": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), - "weight_4": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), - } - } - - RETURN_TYPES = ("EMBEDS",) - FUNCTION = "preprocess" - CATEGORY = "ipadapter" - - def preprocess(self, clip_vision, image_1, ipadapter_plus, noise, weight_1, image_2=None, image_3=None, image_4=None, weight_2=1.0, weight_3=1.0, weight_4=1.0): - weight_1 *= (0.1 + (weight_1 - 0.1)) - weight_1 = 1.19e-05 if weight_1 <= 1.19e-05 else weight_1 - weight_2 *= (0.1 + (weight_2 - 0.1)) - weight_2 = 1.19e-05 if weight_2 <= 1.19e-05 else weight_2 - weight_3 *= (0.1 + (weight_3 - 0.1)) - weight_3 = 1.19e-05 if weight_3 <= 1.19e-05 else weight_3 - weight_4 *= (0.1 + (weight_4 - 0.1)) - weight_5 = 1.19e-05 if weight_4 <= 1.19e-05 else weight_4 - - image = image_1 - weight = [weight_1]*image_1.shape[0] +class ResamplerImport(nn.Module): + def __init__( + self, + dim=1024, + depth=8, + dim_head=64, + heads=16, + num_queries=8, + embedding_dim=768, + output_dim=1024, + ff_mult=4, + ): + super().__init__() - if image_2 is not None: - if image_1.shape[1:] != image_2.shape[1:]: - image_2 = comfy.utils.common_upscale(image_2.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) - image = torch.cat((image, image_2), dim=0) - weight += [weight_2]*image_2.shape[0] - if image_3 is not None: - if image.shape[1:] != image_3.shape[1:]: - image_3 = comfy.utils.common_upscale(image_3.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) - image = torch.cat((image, image_3), dim=0) - weight += [weight_3]*image_3.shape[0] - if image_4 is not None: - if image.shape[1:] != image_4.shape[1:]: - image_4 = comfy.utils.common_upscale(image_4.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) - image = torch.cat((image, image_4), dim=0) - weight += [weight_4]*image_4.shape[0] + self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) - clip_embed = clip_vision.encode_image(image) - neg_image = image_add_noise(image, noise) if noise > 0 else None + self.proj_in = nn.Linear(embedding_dim, dim) + + self.proj_out = nn.Linear(dim, output_dim) + self.norm_out = nn.LayerNorm(output_dim) - if ipadapter_plus: - clip_embed = clip_embed.penultimate_hidden_states - if noise > 0: - clip_embed_zeroed = clip_vision.encode_image(neg_image).penultimate_hidden_states - else: - clip_embed_zeroed = zeroed_hidden_states(clip_vision, image.shape[0]) - else: - clip_embed = clip_embed.image_embeds - if noise > 0: - clip_embed_zeroed = clip_vision.encode_image(neg_image).image_embeds - else: - clip_embed_zeroed = torch.zeros_like(clip_embed) + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append( + nn.ModuleList( + [ + PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), + FeedForward(dim=dim, mult=ff_mult), + ] + ) + ) - if any(e != 1.0 for e in weight): - weight = torch.tensor(weight).unsqueeze(-1) if not ipadapter_plus else torch.tensor(weight).unsqueeze(-1).unsqueeze(-1) - clip_embed = clip_embed * weight + def forward(self, x): - output = torch.stack((clip_embed, clip_embed_zeroed)) - - return( output, ) - -class IPAdapterApplyEncoded(IPAdapterApply): - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "ipadapter": ("IPADAPTER", ), - "embeds": ("EMBEDS",), - "model": ("MODEL", ), - "weight": ("FLOAT", { "default": 1.0, "min": -1, "max": 3, "step": 0.05 }), - "weight_type": (["original", "linear", "channel penalty"], ), - "start_at": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001 }), - "end_at": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001 }), - "unfold_batch": ("BOOLEAN", { "default": False }), - }, - "optional": { - "attn_mask": ("MASK",), - } - } - -class IPAdapterSaveEmbeds: - def __init__(self): - self.output_dir = folder_paths.get_output_directory() - - @classmethod - def INPUT_TYPES(s): - return {"required": { - "embeds": ("EMBEDS",), - "filename_prefix": ("STRING", {"default": "embeds/IPAdapter"}) - }, - } - - RETURN_TYPES = () - FUNCTION = "save" - OUTPUT_NODE = True - CATEGORY = "ipadapter" - - def save(self, embeds, filename_prefix): - full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir) - file = f"{filename}_{counter:05}_.ipadpt" - file = os.path.join(full_output_folder, file) - - torch.save(embeds, file) - return (None, ) - - -class IPAdapterLoadEmbeds: - @classmethod - def INPUT_TYPES(s): - input_dir = folder_paths.get_input_directory() - files = [os.path.relpath(os.path.join(root, file), input_dir) for root, dirs, files in os.walk(input_dir) for file in files if file.endswith('.ipadpt')] - return {"required": {"embeds": [sorted(files), ]}, } - - RETURN_TYPES = ("EMBEDS", ) - FUNCTION = "load" - CATEGORY = "ipadapter" - - def load(self, embeds): - path = folder_paths.get_annotated_filepath(embeds) - output = torch.load(path).cpu() - - return (output, ) - - -class IPAdapterBatchEmbeds: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "embed1": ("EMBEDS",), - "embed2": ("EMBEDS",), - }} - - RETURN_TYPES = ("EMBEDS",) - FUNCTION = "batch" - CATEGORY = "ipadapter" - - def batch(self, embed1, embed2): - output = torch.cat((embed1, embed2), dim=1) - return (output, ) - -NODE_CLASS_MAPPINGS = { - "IPAdapterModelLoader": IPAdapterModelLoader, - "IPAdapterApply": IPAdapterApply, - "IPAdapterApplyEncoded": IPAdapterApplyEncoded, - "IPAdapterEncoder": IPAdapterEncoder, - "IPAdapterSaveEmbeds": IPAdapterSaveEmbeds, - "IPAdapterLoadEmbeds": IPAdapterLoadEmbeds, - "IPAdapterBatchEmbeds": IPAdapterBatchEmbeds, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "IPAdapterModelLoader": "Load IPAdapter Model", - "IPAdapterApply": "Apply IPAdapter", - "IPAdapterApplyEncoded": "Apply IPAdapter from Encoded", - "IPAdapterEncoder": "Encode IPAdapter Image", - "IPAdapterSaveEmbeds": "Save IPAdapter Embeds", - "IPAdapterLoadEmbeds": "Load IPAdapter Embeds", - "IPAdapterBatchEmbeds": "IPAdapter Batch Embeds", -} \ No newline at end of file + latents = self.latents.repeat(x.size(0), 1, 1) + + x = self.proj_in(x) + + for attn, ff in self.layers: + latents = attn(x, latents) + latents + latents = ff(latents) + latents + + latents = self.proj_out(latents) + return self.norm_out(latents)