This commit is contained in:
peter942
2024-01-10 02:21:38 +01:00
parent 1f046a5e15
commit 909566d968
5 changed files with 1664 additions and 1241 deletions
+69 -155
View File
@@ -16,7 +16,11 @@ from .imports.AdvancedControlNet.latent_keyframe_nodes import (
LatentKeyframeInterpolationNodeImport
)
from .imports.AdvancedControlNet.weight_nodes import ScaledSoftUniversalWeightsImport
from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport
from .imports.AdvancedControlNet.control_sparsectrl import SparseIndexMethodImport
from .imports.AdvancedControlNet.nodes import ControlNetLoaderAdvancedImport, AdvancedControlNetApplyImport,TimestepKeyframeNodeImport
from abc import ABC, abstractmethod
class BatchCreativeInterpolationNode:
@classmethod
@@ -27,13 +31,10 @@ class BatchCreativeInterpolationNode:
def INPUT_TYPES(s):
return {
"required": {
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"images": ("IMAGE", ),
"model": ("MODEL", ),
"ipadapter": ("IPADAPTER", ),
"clip_vision": ("CLIP_VISION",),
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
"type_of_frame_distribution": (["linear", "dynamic"],),
"linear_frame_distribution_value": ("INT", {"default": 16, "min": 4, "max": 64, "step": 1}),
"dynamic_frame_distribution_values": ("STRING", {"multiline": True, "default": "0,10,26,40"}),
@@ -41,30 +42,26 @@ class BatchCreativeInterpolationNode:
"linear_key_frame_influence_value": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1}),
"dynamic_key_frame_influence_values": ("STRING", {"multiline": True, "default": "1.0,1.0,1.0,0.5"}),
"type_of_cn_strength_distribution": (["linear", "dynamic"],),
"linear_cn_strength_value": ("STRING", {"multiline": False, "default": "(0.0,0.4)"}),
"linear_cn_strength_value": ("STRING", {"multiline": False, "default": "(0.3,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}),
"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}),
"ipadapter_noise": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
}
}
RETURN_TYPES = ("IMAGE","CONDITIONING","CONDITIONING","MODEL",)
RETURN_NAMES = ("GRAPH","POSITIVE", "NEGATIVE","MODEL")
RETURN_TYPES = ("IMAGE","MODEL","SPARSE_METHOD","INT")
RETURN_NAMES = ("GRAPH","MODEL","KEYFRAME_POSITIONS", "BATCH_SIZE")
FUNCTION = "combined_function"
CATEGORY = "Steerable-Motion/Interpolation"
def combined_function(self, positive, negative, images,model,ipadapter,clip_vision,control_net_name,
type_of_frame_distribution,linear_frame_distribution_value,dynamic_frame_distribution_values,
type_of_key_frame_influence,linear_key_frame_influence_value,dynamic_key_frame_influence_values,
type_of_cn_strength_distribution,linear_cn_strength_value,dynamic_cn_strength_values,
soft_scaled_cn_weights_multiplier,buffer,relative_ipadapter_strength,
relative_ipadapter_influence,ipadapter_noise):
def combined_function(self,images,model,ipadapter,clip_vision,type_of_frame_distribution,
linear_frame_distribution_value, dynamic_frame_distribution_values,
type_of_key_frame_influence,linear_key_frame_influence_value,
dynamic_key_frame_influence_values,type_of_cn_strength_distribution,
linear_cn_strength_value,dynamic_cn_strength_values,buffer,ipadapter_noise):
def calculate_dynamic_influence_ranges(keyframe_positions, key_frame_influence_values, allow_extension=True):
if len(keyframe_positions) < 2 or len(keyframe_positions) != len(key_frame_influence_values):
@@ -74,14 +71,22 @@ class BatchCreativeInterpolationNode:
for i, position in enumerate(keyframe_positions):
influence_factor = key_frame_influence_values[i]
# Calculate the base range size
range_size = influence_factor * (keyframe_positions[-1] - keyframe_positions[0]) / (len(keyframe_positions) - 1) / 2
if i == 0:
# Special handling for the first keyframe (half the distance to the second keyframe)
range_size = influence_factor * (keyframe_positions[1] - keyframe_positions[0]) / 2
start_influence = position # Start from the first keyframe position
end_influence = position + range_size
elif i == len(keyframe_positions) - 1:
# Special handling for the last keyframe (half the distance from the penultimate keyframe)
range_size = influence_factor * (keyframe_positions[-1] - keyframe_positions[-2]) / 2
start_influence = position - range_size
end_influence = position # End at the last keyframe position
else:
# Regular calculation for other keyframes
range_size = influence_factor * (keyframe_positions[-1] - keyframe_positions[0]) / (len(keyframe_positions) - 1) / 2
start_influence = position - range_size
end_influence = position + range_size
# Calculate symmetric start and end influence
start_influence = position - range_size
end_influence = position + range_size
# Adjust start and end influence to not exceed previous and next keyframes
if not allow_extension:
start_influence = max(start_influence, keyframe_positions[i - 1] if i > 0 else 0)
end_influence = min(end_influence, keyframe_positions[i + 1] if i < len(keyframe_positions) - 1 else keyframe_positions[-1])
@@ -90,6 +95,7 @@ class BatchCreativeInterpolationNode:
return influence_ranges
def add_starting_buffer(influence_ranges, buffer=4):
shifted_ranges = [(0, buffer)]
for start, end in influence_ranges:
@@ -126,21 +132,6 @@ class BatchCreativeInterpolationNode:
# Create a list with the linear_key_frame_influence_value for each keyframe
return [linear_key_frame_influence_value for _ in keyframe_positions]
def extract_start_and_endpoint_values(type_of_key_frame_influence, dynamic_key_frame_influence_values, keyframe_positions, linear_key_frame_influence_value):
if type_of_key_frame_influence == "dynamic":
# If dynamic_key_frame_influence_values is a list of characters representing tuples, process it
if isinstance(dynamic_key_frame_influence_values[0], str) and dynamic_key_frame_influence_values[0] == "(":
# Join the characters to form a single string and evaluate to convert into a list of tuples
string_representation = ''.join(dynamic_key_frame_influence_values)
dynamic_values = eval(f'[{string_representation}]')
else:
# If it's already a list of tuples or a single tuple, use it directly
dynamic_values = dynamic_key_frame_influence_values if isinstance(dynamic_key_frame_influence_values, list) else [dynamic_key_frame_influence_values]
return dynamic_values
else:
# Return a list of tuples with the linear_key_frame_influence_value as a tuple repeated for each position
return [linear_key_frame_influence_value for _ in keyframe_positions]
def create_mask_batch(last_key_frame_position, weights, frames):
# Hardcoded dimensions
width, height = 512, 512
@@ -163,67 +154,16 @@ class BatchCreativeInterpolationNode:
return masks_tensor
def adjust_influence_range(batch_index_from, batch_index_to_excl, last_key_frame_position, scale_factor, buffer):
# Calculate the midpoint of the current range
midpoint = (batch_index_from + batch_index_to_excl) // 2
# Calculate the new range length
new_range_length = int((batch_index_to_excl - batch_index_from) * scale_factor)
# Adjusting both sides of the range
if batch_index_from == 0:
# Start is anchored at 0
new_batch_index_from = 0
new_batch_index_to_excl = batch_index_from + new_range_length
elif batch_index_to_excl == last_key_frame_position:
# End is anchored at last_key_frame_position
new_batch_index_from = batch_index_to_excl - new_range_length
new_batch_index_to_excl = last_key_frame_position
else:
# No anchoring, adjust both sides around the midpoint
new_batch_index_from = midpoint - new_range_length // 2
new_batch_index_to_excl = midpoint + new_range_length // 2
# Remove minimum and maximum constraints
return new_batch_index_from, new_batch_index_to_excl
def adjust_strength_values(strength_from, strength_to, multiplier):
mid_point = (strength_from + strength_to) / 2
range_half = abs(strength_to - strength_from) / 2
# Adjust the range with the multiplier
new_range_half = min(range_half * multiplier, 0.5)
# Calculate new strength values, ensuring they stay within [0.0, 1.0]
new_strength_from = max(mid_point - new_range_half, 0.0)
new_strength_to = min(mid_point + new_range_half, 1.0)
# Preserve the order of the original strength values
if strength_from > strength_to:
new_strength_from, new_strength_to = new_strength_to, new_strength_from
return (new_strength_from, new_strength_to)
def plot_weight_comparison(cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights, buffer):
def plot_weight_comparison(ipadapter_frame_numbers, ipadapter_weights, buffer):
plt.figure(figsize=(12, 8))
# Defining colors for each set of data
colors = ['b', 'g', 'r', 'c', 'm', 'y', 'k']
# Alternating the data sets with labels and colors
max_length = max(len(cn_frame_numbers), len(ipadapter_frame_numbers))
# Plotting data for ipadapter
max_length = len(ipadapter_frame_numbers)
label_counter = 1 if buffer < 0 else 0 # Start from 1 if buffer < 0, else start from 0
for i in range(max_length):
# Label for cn_strength
if i < len(cn_frame_numbers):
if i == 0 and buffer > 0:
label = 'cn_strength_buffer'
else:
label = f'cn_strength_{label_counter}'
plt.plot(cn_frame_numbers[i], cn_weights[i], marker='o', color=colors[i % len(colors)], label=label)
# Label for ipa_strength
if i < len(ipadapter_frame_numbers):
if i == 0 and buffer > 0:
label = 'ipa_strength_buffer'
@@ -235,7 +175,7 @@ class BatchCreativeInterpolationNode:
label_counter += 1
plt.legend()
max_weight = max([weight.max() for weight in cn_weights + ipadapter_weights]) * 1.5
max_weight = max([weight.max() for weight in ipadapter_weights]) * 1.5
plt.ylim(0, max_weight)
buffer_io = BytesIO()
@@ -253,25 +193,41 @@ class BatchCreativeInterpolationNode:
return (img_tensor,)
def extract_start_and_endpoint_values(type_of_key_frame_influence, dynamic_key_frame_influence_values, keyframe_positions, linear_key_frame_influence_value):
if type_of_key_frame_influence == "dynamic":
# If dynamic_key_frame_influence_values is a list of characters representing tuples, process it
if isinstance(dynamic_key_frame_influence_values[0], str) and dynamic_key_frame_influence_values[0] == "(":
# Join the characters to form a single string and evaluate to convert into a list of tuples
string_representation = ''.join(dynamic_key_frame_influence_values)
dynamic_values = eval(f'[{string_representation}]')
else:
# If it's already a list of tuples or a single tuple, use it directly
dynamic_values = dynamic_key_frame_influence_values if isinstance(dynamic_key_frame_influence_values, list) else [dynamic_key_frame_influence_values]
return dynamic_values
else:
# Return a list of tuples with the linear_key_frame_influence_value as a tuple repeated for each position
return [linear_key_frame_influence_value for _ in keyframe_positions]
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 = 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 = [], [], [], []
keyframe_positions_string = ','.join(str(pos) for pos in keyframe_positions)
sparseindexmethod = SparseIndexMethodNodeImport()
sparse_indexes, = sparseindexmethod.get_method(keyframe_positions_string)
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 = add_starting_buffer(influence_ranges, buffer)
last_key_frame_position = (keyframe_positions[-1]) + buffer
all_frame_numbers = []
all_weights = []
embeds = []
masks = []
existing_embeds = []
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
for i, (batch_index_from, batch_index_to_excl) in enumerate(influence_ranges):
# Default values
revert_direction_at_midpoint = False
@@ -282,85 +238,43 @@ class BatchCreativeInterpolationNode:
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)
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)
interpolation = "ease-in"
# 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)
interpolation = "ease-out"
# 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)
revert_direction_at_midpoint = True
# Import necessary modules
latent_keyframe_interpolation_node = LatentKeyframeInterpolationNodeImport()
scaled_soft_control_net_weights = ScaledSoftUniversalWeightsImport()
timestep_keyframe_node = TimestepKeyframeNodeImport()
control_net_loader = ControlNetLoaderAdvancedImport()
apply_advanced_control_net = AdvancedControlNetApplyImport()
ipadapter_application = IPAdapterApplyImport()
ipadapter_encoder = IPAdapterEncoderImport()
# ipadapter_batcher = IPAdapterBatchEmbedsImport()
# 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)
# 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, latent_keyframe=latent_keyframe, prev_timestep_keyframe=None)[0]
# 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)
# Prepare image
prepped_image = prep_image(image=image.unsqueeze(0), interpolation="LANCZOS", crop_position="pad", sharpening=0.0)[0]
# 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)
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, 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)
mask = create_mask_batch(last_key_frame_position, weights, frame_numbers)
mask = create_mask_batch(last_key_frame_position, ipa_weights, frame_numbers)
# add mask to masks list
masks.append(mask)
embed, = ipadapter_encoder.preprocess(clip_vision, prepped_image, True, 0.0, 1.0)
# add embeds to current batch
embeds.append(embed)
model, = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, image=None, weight_type="original",
noise=ipadapter_noise, embeds=embed, attn_mask=mask, start_at=0.0, end_at=1.0, unfold_batch=True)
all_frame_numbers.append(frame_numbers)
all_weights.append(weights)
# print out the format for the embeds
weights_diagram, = plot_weight_comparison(all_frame_numbers, all_weights, buffer)
# merged_embeds = torch.cat(embeds, dim=1)
# stacked_masks = torch.stack(masks)
# merged_masks = torch.cat(masks, dim=1)
comparison_diagram, = plot_weight_comparison(cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights, buffer)
return comparison_diagram, positive, negative, model
return weights_diagram, model,sparse_indexes, last_key_frame_position
# NODE MAPPING
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,79 @@
#taken from: https://github.com/lllyasviel/ControlNet
#and modified
#and then taken from comfy/cldm/cldm.py and modified again
from abc import ABC, abstractmethod
import math
import numpy as np
from typing import Iterable, Union
import torch
import torch as th
import torch.nn as nn
from torch import Tensor
from einops import rearrange, repeat
from comfy.ldm.modules.diffusionmodules.util import (
zero_module,
timestep_embedding,
)
from comfy.cldm.cldm import ControlNet as ControlNetCLDM
from comfy.ldm.modules.attention import SpatialTransformer
from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample
from comfy.ldm.util import exists
from comfy.ldm.modules.attention import default, optimized_attention
from comfy.ldm.modules.attention import FeedForward, SpatialTransformer
from comfy.controlnet import broadcast_image_to
from comfy.utils import repeat_to_batch_size
import comfy.ops
# from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch
class SparseMethodImport(ABC):
SPREAD = "spread"
INDEX = "index"
def __init__(self, method: str):
self.method = method
@abstractmethod
def get_indexes(self, hint_length: int, full_length: int) -> list[int]:
pass
class SparseIndexMethodImport(SparseMethodImport):
def __init__(self, idxs: list[int]):
super().__init__(self.INDEX)
self.idxs = idxs
def get_indexes(self, hint_length: int, full_length: int) -> list[int]:
orig_hint_length = hint_length
if hint_length > full_length:
hint_length = full_length
# if idxs is less than hint_length, throw error
if len(self.idxs) < hint_length:
err_msg = f"There are not enough indexes ({len(self.idxs)}) provided to fit the usable {hint_length} input images."
if orig_hint_length != hint_length:
err_msg = f"{err_msg} (original input images: {orig_hint_length})"
raise ValueError(err_msg)
# cap idxs to hint_length
idxs = self.idxs[:hint_length]
new_idxs = []
real_idxs = set()
for idx in idxs:
if idx < 0:
real_idx = full_length+idx
if real_idx in real_idxs:
raise ValueError(f"Index '{idx}' maps to '{real_idx}' and is duplicate - indexes in Sparse Index Method must be unique.")
else:
real_idx = idx
if real_idx in real_idxs:
raise ValueError(f"Index '{idx}' is duplicate (or a negative index is equivalent) - indexes in Sparse Index Method must be unique.")
real_idxs.add(real_idx)
new_idxs.append(real_idx)
return new_idxs
@@ -248,10 +248,18 @@ def calculate_weights(batch_index_from, batch_index_to, strength_from, strength_
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]])
'''
peak_reduction = 2
if peak_reduction > 0:
mid_point = len(weights) // 2
start = mid_point - peak_reduction // 2
end = mid_point + peak_reduction // 2
weights = np.concatenate([weights[:start], weights[end:]])
'''
# Generate frame numbers
frame_numbers = np.arange(range_start, range_start + len(weights))
@@ -0,0 +1,44 @@
from torch import Tensor
import folder_paths
from nodes import VAEEncode
import comfy.utils
# from .utils import TimestepKeyframeGroup
from .control_sparsectrl import SparseIndexMethodImport
# from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced
class SparseIndexMethodNodeImport:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"indexes": ("STRING", {"default": "0"}),
}
}
RETURN_TYPES = ("SPARSE_METHOD",)
FUNCTION = "get_method"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
def get_method(self, indexes: str):
idxs = []
unique_idxs = set()
# get indeces from string
str_idxs = [x.strip() for x in indexes.strip().split(",")]
for str_idx in str_idxs:
try:
idx = int(str_idx)
if idx in unique_idxs:
raise ValueError(f"'{idx}' is duplicated; indexes must be unique.")
idxs.append(idx)
unique_idxs.add(idx)
except ValueError:
raise ValueError(f"'{str_idx}' is not a valid integer index.")
if len(idxs) == 0:
raise ValueError(f"No indexes were listed in Sparse Index Method.")
return (SparseIndexMethodImport(idxs),)