1.1
This commit is contained in:
+69
-155
@@ -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
|
||||
|
||||
+1442
-1064
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),)
|
||||
|
||||
Reference in New Issue
Block a user