From 95a894c24291c6bef2f3f57856c006d8790c29ae Mon Sep 17 00:00:00 2001 From: peter942 Date: Tue, 5 Dec 2023 17:34:12 +0100 Subject: [PATCH] =?UTF-8?q?=F0=9F=8E=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.yaml | 3 + control/.DS_Store | Bin 0 -> 6148 bytes control/film.py | 201 ++++++++++++++++++++++++++++++++++++++++++++++ control/nodes.py | 79 +++++++++++------- 4 files changed, 253 insertions(+), 30 deletions(-) create mode 100644 config.yaml create mode 100644 control/.DS_Store create mode 100644 control/film.py diff --git a/config.yaml b/config.yaml new file mode 100644 index 0000000..b99d4a7 --- /dev/null +++ b/config.yaml @@ -0,0 +1,3 @@ +#Plz don't delete this file, just edit it when neccessary. +ckpts_path: "./ckpts" +ops_backend: "cupy" #Either "taichi" or "cupy" \ No newline at end of file diff --git a/control/.DS_Store b/control/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..5008ddfcf53c02e82d7eee2e57c38e5672ef89f6 GIT binary patch literal 6148 zcmeH~Jr2S!425mzP>H1@V-^m;4Wg<&0T*E43hX&L&p$$qDprKhvt+--jT7}7np#A3 zem<@ulZcFPQ@L2!n>{z**++&mCkOWA81W14cNZlEfg7;MkzE(HCqgga^y>{tEnwC%0;vJ&^%eQ zLs35+`xjp>T0 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), ) \ No newline at end of file diff --git a/control/nodes.py b/control/nodes.py index 31b1275..7ffafcc 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -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 🛂🅐🅒🅝",