Refactoring
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
@@ -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']
|
||||
|
||||
Vendored
BIN
Binary file not shown.
@@ -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)
|
||||
@@ -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,)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user