🎨
This commit is contained in:
Vendored
BIN
Binary file not shown.
+201
@@ -0,0 +1,201 @@
|
||||
import torch
|
||||
from comfy.model_management import get_torch_device, soft_empty_cache
|
||||
import bisect
|
||||
from typing import List
|
||||
import numpy as np
|
||||
import ast
|
||||
import typing
|
||||
import pathlib
|
||||
import einops
|
||||
import traceback
|
||||
import os
|
||||
from urllib.parse import urlparse
|
||||
from torch.hub import download_url_to_file, get_dir
|
||||
import yaml
|
||||
|
||||
|
||||
BASE_MODEL_DOWNLOAD_URLS = [
|
||||
"https://github.com/styler00dollar/VSGAN-tensorrt-docker/releases/download/models/",
|
||||
"https://github.com/Fannovel16/ComfyUI-Frame-Interpolation/releases/download/models/",
|
||||
"https://github.com/dajes/frame-interpolation-pytorch/releases/download/v1.0.0/"
|
||||
]
|
||||
|
||||
config_path = os.path.join(os.path.dirname(__file__), "./config.yaml")
|
||||
if os.path.exists(config_path):
|
||||
config = yaml.load(open(config_path, "r"), Loader=yaml.FullLoader)
|
||||
else:
|
||||
raise Exception("config.yaml file is neccessary, plz recreate the config file by downloading it from https://github.com/Fannovel16/ComfyUI-Frame-Interpolation")
|
||||
DEVICE = get_torch_device()
|
||||
|
||||
def get_ckpt_container_path(model_type):
|
||||
return os.path.abspath(os.path.join(os.path.dirname(__file__), config["ckpts_path"], model_type))
|
||||
|
||||
|
||||
def load_file_from_url(url, model_dir=None, progress=True, file_name=None):
|
||||
|
||||
if model_dir is None: # use the pytorch hub_dir
|
||||
hub_dir = get_dir()
|
||||
model_dir = os.path.join(hub_dir, 'checkpoints')
|
||||
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
|
||||
parts = urlparse(url)
|
||||
file_name = os.path.basename(parts.path)
|
||||
if file_name is not None:
|
||||
file_name = file_name
|
||||
cached_file = os.path.abspath(os.path.join(model_dir, file_name))
|
||||
if not os.path.exists(cached_file):
|
||||
print(f'Downloading: "{url}" to {cached_file}\n')
|
||||
download_url_to_file(url, cached_file, hash_prefix=None, progress=progress)
|
||||
return cached_file
|
||||
|
||||
class InterpolationStateList():
|
||||
|
||||
def __init__(self, frame_indices: typing.List[int], is_skip_list: bool):
|
||||
self.frame_indices = frame_indices
|
||||
self.is_skip_list = is_skip_list
|
||||
|
||||
def is_frame_skipped(self, frame_index):
|
||||
is_frame_in_list = frame_index in self.frame_indices
|
||||
return self.is_skip_list and is_frame_in_list or not self.is_skip_list and not is_frame_in_list
|
||||
|
||||
|
||||
def load_file_from_github_release(model_type, ckpt_name):
|
||||
error_strs = []
|
||||
for i, base_model_download_url in enumerate(BASE_MODEL_DOWNLOAD_URLS):
|
||||
try:
|
||||
return load_file_from_url(base_model_download_url + ckpt_name, get_ckpt_container_path(model_type))
|
||||
except Exception:
|
||||
traceback_str = traceback.format_exc()
|
||||
if i < len(BASE_MODEL_DOWNLOAD_URLS) - 1:
|
||||
print("Failed! Trying another endpoint.")
|
||||
error_strs.append(f"Error when downloading from: {base_model_download_url + ckpt_name}\n\n{traceback_str}")
|
||||
|
||||
error_str = '\n\n'.join(error_strs)
|
||||
raise Exception(f"Tried all GitHub base urls to download {ckpt_name} but no suceess. Below is the error log:\n\n{error_str}")
|
||||
|
||||
|
||||
def preprocess_frames(frames):
|
||||
return einops.rearrange(frames, "n h w c -> n c h w")
|
||||
|
||||
def postprocess_frames(frames):
|
||||
return einops.rearrange(frames, "n c h w -> n h w c").cpu()
|
||||
|
||||
|
||||
MODEL_TYPE = pathlib.Path(__file__).parent.name
|
||||
DEVICE = get_torch_device()
|
||||
|
||||
def inference(model, img_batch_1, img_batch_2, inter_frames):
|
||||
results = [
|
||||
img_batch_1,
|
||||
img_batch_2
|
||||
]
|
||||
|
||||
idxes = [0, inter_frames + 1]
|
||||
remains = list(range(1, inter_frames + 1))
|
||||
|
||||
splits = torch.linspace(0, 1, inter_frames + 2)
|
||||
|
||||
for _ in range(len(remains)):
|
||||
starts = splits[idxes[:-1]]
|
||||
ends = splits[idxes[1:]]
|
||||
distances = ((splits[None, remains] - starts[:, None]) / (ends[:, None] - starts[:, None]) - .5).abs()
|
||||
matrix = torch.argmin(distances).item()
|
||||
start_i, step = np.unravel_index(matrix, distances.shape)
|
||||
end_i = start_i + 1
|
||||
|
||||
x0 = results[start_i].to(DEVICE)
|
||||
x1 = results[end_i].to(DEVICE)
|
||||
dt = x0.new_full((1, 1), (splits[remains[step]] - splits[idxes[start_i]])) / (splits[idxes[end_i]] - splits[idxes[start_i]])
|
||||
|
||||
with torch.no_grad():
|
||||
prediction = model(x0, x1, dt)
|
||||
insert_position = bisect.bisect_left(idxes, remains[step])
|
||||
idxes.insert(insert_position, remains[step])
|
||||
results.insert(insert_position, prediction.clamp(0, 1).float())
|
||||
del remains[step]
|
||||
|
||||
return [tensor.flip(0) for tensor in results]
|
||||
|
||||
|
||||
def film_interpolation(
|
||||
frames: torch.Tensor,
|
||||
frame_counts: List[int] = [0,5,32,40],
|
||||
buffer: int = 10):
|
||||
|
||||
frame_counts = sorted(frame_counts)
|
||||
|
||||
max_gap = max(b-a for a, b in zip(frame_counts[:-1], frame_counts[1:]))
|
||||
|
||||
model_path = load_file_from_github_release(MODEL_TYPE, "film_net_fp32.pt")
|
||||
model = torch.jit.load(model_path, map_location='cpu')
|
||||
model.eval()
|
||||
model = model.to(DEVICE)
|
||||
|
||||
frames = preprocess_frames(frames)
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache = 0
|
||||
clear_cache_after_n_frames = 10 # Example value
|
||||
|
||||
# Generate buffer frames and attach them to the beginning of output_frames
|
||||
first_frame = frames[0].unsqueeze(0)
|
||||
buffer_frames = [first_frame] * buffer
|
||||
output_frames = buffer_frames
|
||||
|
||||
|
||||
for frame_itr in range(len(frames) - 1):
|
||||
|
||||
|
||||
frame_0 = frames[frame_itr:frame_itr+1].to(DEVICE)
|
||||
frame_1 = frames[frame_itr+1:frame_itr+2].to(DEVICE)
|
||||
frame_output = []
|
||||
result = inference(model, frame_0, frame_1, max_gap - 1)
|
||||
|
||||
# Find the current frame's position in the frame_counts list
|
||||
current_frame = frame_counts[frame_itr]
|
||||
next_frame = frame_counts[frame_itr + 1] - 1
|
||||
current_gap = next_frame - current_frame
|
||||
# Determine the number of frames to drop based on the difference between the max gap and the current gap
|
||||
frames_to_drop = max_gap - current_gap
|
||||
|
||||
frame_output = result[:-1]
|
||||
|
||||
# frames_to_drop = max_gap - 1 - len(frame_output)
|
||||
if frames_to_drop > 0:
|
||||
if frames_to_drop >= len(frame_output):
|
||||
raise ValueError("Number of frames to drop is greater than or equal to total number of frames in the batch.")
|
||||
|
||||
drop_interval = len(frame_output) / float(frames_to_drop)
|
||||
result = []
|
||||
next_drop = drop_interval
|
||||
|
||||
for i, frame in enumerate(frame_output):
|
||||
if i >= next_drop:
|
||||
next_drop += drop_interval
|
||||
else:
|
||||
result.append(frame)
|
||||
|
||||
frame_output = result # Update frame_output with only the undropped frames
|
||||
|
||||
# Detach and move to CPU all frames, whether or not any were dropped
|
||||
frame_output = [frame.detach().cpu() for frame in frame_output]
|
||||
|
||||
# Append the processed (and potentially dropped) frames to the main list
|
||||
output_frames.extend(frame_output)
|
||||
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache += 1
|
||||
if number_of_frames_processed_since_last_cleared_cuda_cache >= clear_cache_after_n_frames:
|
||||
print("Comfy-VFI: Clearing cache...")
|
||||
soft_empty_cache()
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache = 0
|
||||
print("Comfy-VFI: Done cache clearing")
|
||||
|
||||
output_frames.append(frames[-1:])
|
||||
out = torch.cat(output_frames, dim=0)
|
||||
|
||||
# clear cache for courtesy
|
||||
print("Comfy-VFI: Final clearing cache...")
|
||||
soft_empty_cache()
|
||||
print("Comfy-VFI: Done cache clearing")
|
||||
|
||||
|
||||
return (postprocess_frames(out), )
|
||||
+49
-30
@@ -1,14 +1,16 @@
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import folder_paths
|
||||
from ast import literal_eval
|
||||
from .control import ControlNetAdvancedImport, T2IAdapterAdvancedImport, load_controlnet, ControlNetWeightsTypeImport, T2IAdapterWeightsTypeImport,\
|
||||
LatentKeyframeGroupImport, TimestepKeyframeImport, TimestepKeyframeGroupImport, is_advanced_controlnet
|
||||
|
||||
from .weight_nodes import ScaledSoftControlNetWeightsImport, SoftControlNetWeightsImport, CustomControlNetWeightsImport, \
|
||||
SoftT2IAdapterWeightsImport, CustomT2IAdapterWeightsImport
|
||||
from .latent_keyframe_nodes import LatentKeyframeGroupNodeImport, LatentKeyframeInterpolationNodeImport, LatentKeyframeBatchedGroupNodeImport, LatentKeyframeNodeImport
|
||||
from .deprecated_nodes import LoadImagesFromDirectory
|
||||
from .logger import logger
|
||||
from .film import film_interpolation
|
||||
|
||||
|
||||
class TimestepKeyframeNodeImport:
|
||||
@@ -151,6 +153,37 @@ class AdvancedControlNetApplyImport:
|
||||
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
|
||||
@@ -177,18 +210,19 @@ class BatchCreativeInterpolationNode:
|
||||
"soft_scaled_cn_weights_multiplier": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"interpolation": (["ease-in", "ease-out", "ease-in-out"],),
|
||||
"buffer": ("INT", {"default": 4, "min": 1, "max": 16, "step": 1}),
|
||||
"intermediate_frame_mask_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
"optional": {
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING","CONDITIONING")
|
||||
RETURN_TYPES = ("CONDITIONING","CONDITIONING","IMAGE")
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
FUNCTION = "combined_function"
|
||||
|
||||
CATEGORY = "ComfyUI-Creative-Interpolation 🎞️🅟🅞🅜/Interpolation"
|
||||
CATEGORY = "Steerable-Motion/Interpolation"
|
||||
|
||||
def combined_function(self, positive, negative, control_net_name, images,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,interpolation,buffer):
|
||||
def combined_function(self, positive, negative, control_net_name, images,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,interpolation,buffer,intermediate_frame_mask_strength):
|
||||
|
||||
def calculate_dynamic_influence_ranges(keyframe_positions, key_frame_influence_values):
|
||||
if len(keyframe_positions) < 2 or len(keyframe_positions) != len(key_frame_influence_values):
|
||||
@@ -248,7 +282,6 @@ class BatchCreativeInterpolationNode:
|
||||
else:
|
||||
# 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":
|
||||
@@ -264,39 +297,23 @@ class BatchCreativeInterpolationNode:
|
||||
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]
|
||||
|
||||
print("type_of_frame_distribution",type_of_frame_distribution)
|
||||
print("dynamic_frame_distribution_values",dynamic_frame_distribution_values)
|
||||
print("linear_frame_distribution_value",linear_frame_distribution_value)
|
||||
print("type_of_key_frame_influence",type_of_key_frame_influence)
|
||||
print("linear_key_frame_influence_value",linear_key_frame_influence_value)
|
||||
print("dynamic_key_frame_influence_values",dynamic_key_frame_influence_values)
|
||||
print("type_of_cn_strength_distribution",type_of_cn_strength_distribution)
|
||||
print("linear_cn_strength_value",linear_cn_strength_value)
|
||||
print("dynamic_cn_strength_values",dynamic_cn_strength_values)
|
||||
print("soft_scaled_cn_weights_multiplier",soft_scaled_cn_weights_multiplier)
|
||||
print("interpolation",interpolation)
|
||||
print("buffer",buffer)
|
||||
|
||||
|
||||
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)
|
||||
if buffer > 0:
|
||||
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]
|
||||
|
||||
print("keyframe_positions",keyframe_positions)
|
||||
print("cn_strength_values",cn_strength_values)
|
||||
print("key_frame_influence_values",key_frame_influence_values)
|
||||
print("influence_ranges",influence_ranges)
|
||||
|
||||
ipadapter_input, = film_interpolation(images, keyframe_positions, buffer)
|
||||
|
||||
last_key_frame_position = (keyframe_positions[-1]) + buffer
|
||||
control_net = []
|
||||
for i, (start, end) in enumerate(influence_ranges):
|
||||
batch_index_from, batch_index_to_excl = influence_ranges[i]
|
||||
|
||||
if i == 0: # buffer image
|
||||
if i == 0 and 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)
|
||||
return_at_midpoint = False
|
||||
@@ -313,7 +330,6 @@ class BatchCreativeInterpolationNode:
|
||||
strength_from, strength_to = cn_strength_values[i-1] if i-1 < len(cn_strength_values) else (0.0, 1.0)
|
||||
return_at_midpoint = True
|
||||
|
||||
|
||||
latent_keyframe_interpolation_node = LatentKeyframeInterpolationNodeImport()
|
||||
latent_keyframe, = latent_keyframe_interpolation_node.load_keyframe(
|
||||
batch_index_from,
|
||||
@@ -354,12 +370,14 @@ class BatchCreativeInterpolationNode:
|
||||
0.0,
|
||||
1.0)
|
||||
|
||||
return positive, negative
|
||||
return positive, negative, ipadapter_input
|
||||
|
||||
# NODE MAPPING
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# Combined
|
||||
"BatchCreativeInterpolation": BatchCreativeInterpolationNode
|
||||
"BatchCreativeInterpolation": BatchCreativeInterpolationNode,
|
||||
"MaskGenerator": MaskGeneratorNode
|
||||
# "FILMVFIImport": FILMVFINode
|
||||
# Keyframes
|
||||
# "TimestepKeyframe": TimestepKeyframeNodeImport,
|
||||
# "LatentKeyframeImport": LatentKeyframeNodeImport,
|
||||
@@ -383,7 +401,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Combined
|
||||
"BatchCreativeInterpolation": "Batch Creative Interpolation 🎞️🅟🅞🅜"
|
||||
"BatchCreativeInterpolation": "Batch Creative Interpolation 🎞️🅟🅞🅜",
|
||||
"MaskGenerator": "Mask Generator 🎞️🅟🅞🅜"
|
||||
# Keyframes
|
||||
# "TimestepKeyframe": "Timestep Keyframe 🎞️🅟🅞🅜",
|
||||
# "LatentKeyframe": "Latent Keyframe 🛂🅐🅒🅝",
|
||||
|
||||
Reference in New Issue
Block a user