Refactoring

This commit is contained in:
peter942
2023-12-07 16:16:13 +01:00
parent 023c75281a
commit 5570e3c3fd
12 changed files with 618 additions and 1220 deletions
Vendored
BIN
View File
Binary file not shown.
+6
View File
@@ -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:
+56 -291
View File
@@ -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 🎞️🅢🅜"
}
+1 -1
View File
@@ -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']
BIN
View File
Binary file not shown.
-70
View File
@@ -1,70 +0,0 @@
import os
import torch
import numpy as np
from PIL import Image, ImageOps
from .logger import logger
class LoadImagesFromDirectory:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"directory": ("STRING", {"default": ""}),
},
"optional": {
"image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}),
"start_index": ("INT", {"default": 0, "min": 0, "step": 1}),
}
}
RETURN_TYPES = ("IMAGE", "MASK", "INT")
FUNCTION = "load_images"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/deprecated"
def load_images(self, directory: str, image_load_cap: int = 0, start_index: int = 0):
if not os.path.isdir(directory):
raise FileNotFoundError(f"Directory '{directory} cannot be found.'")
dir_files = os.listdir(directory)
if len(dir_files) == 0:
raise FileNotFoundError(f"No files in directory '{directory}'.")
dir_files = sorted(dir_files)
dir_files = [os.path.join(directory, x) for x in dir_files]
# start at start_index
dir_files = dir_files[start_index:]
images = []
masks = []
limit_images = False
if image_load_cap > 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)
-296
View File
@@ -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,)
-36
View File
@@ -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)
-121
View File
@@ -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)
-157
View File
@@ -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)))
@@ -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
@@ -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",
}
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)