This commit is contained in:
peter942
2023-12-05 17:34:12 +01:00
parent 2a584546b7
commit 95a894c242
4 changed files with 253 additions and 30 deletions
BIN
View File
Binary file not shown.
+201
View File
@@ -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
View File
@@ -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 🛂🅐🅒🅝",