diff --git a/SteerableMotion.py b/SteerableMotion.py index 5afd57d..2cf7868 100644 --- a/SteerableMotion.py +++ b/SteerableMotion.py @@ -10,7 +10,9 @@ import matplotlib.pyplot as plt # Local application/library specific imports from .imports.ComfyUI_IPAdapter_plus.IPAdapterPlus import IPAdapterBatchImport, IPAdapterTiledBatchImport, IPAdapterTiledImport, PrepImageForClipVisionImport, IPAdapterAdvancedImport, IPAdapterNoiseImport from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport +from .imports.ComfyUI_Frame_Interpolation.vfi_models.film import FILM_VFIImport import matplotlib +import gc class BatchCreativeInterpolationNode: @classmethod @@ -563,8 +565,12 @@ class BatchCreativeInterpolationNode: return comparison_diagram, positive, negative, model, sparse_indexes, last_key_frame_position, buffer, shifted_keyframes_position -class DropFramesByIndex: +# import the class FILM_VFI from ComfyUI-Frame-Interpolation/vfi_models/film/__init__.py +# from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport + + +class RemoveAndInterpolateFramesNode: @classmethod def INPUT_TYPES(s): return { @@ -577,26 +583,34 @@ class DropFramesByIndex: RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("image",) - FUNCTION = "drop_frames_by_index" - + FUNCTION = "replace_and_interpolate_frames" CATEGORY = "Steerable-Motion" - def drop_frames_by_index(self, images: torch.Tensor, frames_to_drop: str): - # Convert the string of frame indices to a list of integers - print(frames_to_drop) - print(type(frames_to_drop)) - + def replace_and_interpolate_frames(self, images: torch.Tensor, frames_to_drop: str): if isinstance(frames_to_drop, str): frames_to_drop = eval(frames_to_drop) - # Sort and reverse the list of frames to drop to avoid index out of range error when removing frames_to_drop = sorted(frames_to_drop, reverse=True) - - # Drop the frames by index - for index in frames_to_drop: - if index < images.shape[0]: - images = torch.cat((images[:index], images[index+1:])) + # Create instance of FILM_VFI within the function + film_vfi = FILM_VFIImport() # Assuming FILM_VFI does not require any special setup + + for index in frames_to_drop: + if 0 < index < images.shape[0] - 1: + # Extract the two surrounding frames + batch = images[index-1:index+2:2] + + # Process through FILM_VFI + interpolated_frames = film_vfi.vfi( + ckpt_name='film_net_fp32.pt', + frames=batch, + clear_cache_after_n_frames=10, + multiplier=2 + )[0] # Assuming vfi returns a tuple and the first element is the interpolated frames + + # Replace the original frames at the location + images = torch.cat((images[:index-1], interpolated_frames, images[index+2:])) + return (images,) @@ -644,11 +658,11 @@ class IpaConfigurationNode: NODE_CLASS_MAPPINGS = { "BatchCreativeInterpolation": BatchCreativeInterpolationNode, "IpaConfiguration": IpaConfigurationNode, - "DropFramesByIndex": DropFramesByIndex, + "RemoveAndInterpolateFrames": RemoveAndInterpolateFramesNode, } NODE_DISPLAY_NAME_MAPPINGS = { "BatchCreativeInterpolation": "Batch Creative Interpolation 🎞️🅢🅜", "IpaConfiguration": "IPA Configuration 🎞️🅢🅜", - "DropFramesByIndex": "Drop Frames By Index 🎞️🅢🅜", + "RemoveAndInterpolateFrames": "Remove and Interpolate Frames 🎞️🅢🅜", } \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/.gitignore b/imports/ComfyUI_Frame_Interpolation/.gitignore new file mode 100644 index 0000000..c16001c --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/.gitignore @@ -0,0 +1,3 @@ +ckpts +__pycache__ +test_result \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/All_in_one_v1_3.png b/imports/ComfyUI_Frame_Interpolation/All_in_one_v1_3.png new file mode 100644 index 0000000..091ccfa Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/All_in_one_v1_3.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/LICENSE b/imports/ComfyUI_Frame_Interpolation/LICENSE new file mode 100644 index 0000000..2a8000a --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2023 Fannovel16 + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/README.md b/imports/ComfyUI_Frame_Interpolation/README.md new file mode 100644 index 0000000..f5d0f3c --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/README.md @@ -0,0 +1,194 @@ +# ComfyUI Frame Interpolation (ComfyUI VFI) (WIP) + +A custom node set for Video Frame Interpolation in ComfyUI. +**UPDATE** Memory management is improved. Now this extension takes less RAM and VRAM than before. + +**UPDATE 2** VFI nodes now accept scheduling multipiler values + +![](./interpolation_schedule.png) +![](./test_vfi_schedule.gif) + +## Nodes +* KSampler Gradually Adding More Denoise (efficient) +* GMFSS Fortuna VFI +* IFRNet VFI +* IFUnet VFI +* M2M VFI +* RIFE VFI (4.0 - 4.9) (Note that option `fast_mode` won't do anything from v4.5+ as `contextnet` is removed) +* FILM VFI +* Sepconv VFI +* AMT VFI +* Make Interpolation State List +* STMFNet VFI (requires at least 4 frames, can only do 2x interpolation for now) +* FLAVR VFI (same conditions as STMFNet) + +## Install +### ComfyUI Manager +Incompatibile issue with it is now fixed + +Following this guide to install this extension + +https://github.com/ltdrdata/ComfyUI-Manager#how-to-use +### Command-line +#### Windows +Run install.bat + +For Window users, if you are having trouble with cupy, please run `install.bat` instead of `install-cupy.py` or `python install.py`. +#### Linux +Open your shell app and start venv if it is used for ComfyUI. Run: +``` +python install.py +``` +## Support for non-CUDA device (experimental) +If you don't have a NVidia card, you can try `taichi` ops backend powered by [Taichi Lang](https://www.taichi-lang.org/) + +On Windows, you can install it by running `install.bat` or `pip install taichi` on Linux + +Then change value of `ops_backend` from `cupy` to `taichi` in `config.yaml` + +If `NotImplementedError` appears, a VFI node in the workflow isn't supported by taichi + +## Usage +All VFI nodes can be accessed in **category** `ComfyUI-Frame-Interpolation/VFI` if the installation is successful and require a `IMAGE` containing frames (at least 2, or at least 4 for STMF-Net/FLAVR). + +Regarding STMFNet and FLAVR, if you only have two or three frames, you should use: Load Images -> Other VFI node (FILM is recommended in this case) with `multiplier=4` -> STMFNet VFI/FLAVR VFI + +`clear_cache_after_n_frames` is used to avoid out-of-memory. Decreasing it makes the chance lower but also increases processing time. + +It is recommended to use LoadImages (LoadImagesFromDirectory) from [ComfyUI-Advanced-ControlNet](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/) and [ComfyUI-VideoHelperSuite](https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite) along side with this extension. + +## Example +### Simple workflow +Workflow metadata isn't embeded +Download these two images [anime0.png](./demo_frames/anime0.png) and [anime1.png](./demo_frames/anime0.png) and put them into a folder like `E:\test` in this image. +![](./example.png) + +### Complex workflow +It's used in AnimationDiff (can load workflow metadata) +![](All_in_one_v1_3.png) + +## Credit +Big thanks for styler00dollar for making [VSGAN-tensorrt-docker](https://github.com/styler00dollar/VSGAN-tensorrt-docker). About 99% the code of this repo comes from it. + +Citation for each VFI node: +### GMFSS Fortuna +The All-In-One GMFSS: Dedicated for Anime Video Frame Interpolation + +https://github.com/98mxr/GMFSS_Fortuna + +### IFRNet +```bibtex +@InProceedings{Kong_2022_CVPR, + author = {Kong, Lingtong and Jiang, Boyuan and Luo, Donghao and Chu, Wenqing and Huang, Xiaoming and Tai, Ying and Wang, Chengjie and Yang, Jie}, + title = {IFRNet: Intermediate Feature Refine Network for Efficient Frame Interpolation}, + booktitle = {Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)}, + year = {2022} +} +``` + +### IFUnet +RIFE with IFUNet, FusionNet and RefineNet + +https://github.com/98mxr/IFUNet +### M2M +```bibtex +@InProceedings{hu2022m2m, + title={Many-to-many Splatting for Efficient Video Frame Interpolation}, + author={Hu, Ping and Niklaus, Simon and Sclaroff, Stan and Saenko, Kate}, + journal={CVPR}, + year={2022} + } +``` + +### RIFE +```bibtex +@inproceedings{huang2022rife, + title={Real-Time Intermediate Flow Estimation for Video Frame Interpolation}, + author={Huang, Zhewei and Zhang, Tianyuan and Heng, Wen and Shi, Boxin and Zhou, Shuchang}, + booktitle={Proceedings of the European Conference on Computer Vision (ECCV)}, + year={2022} +} +``` + +### FILM +[Frame interpolation in PyTorch](https://github.com/dajes/frame-interpolation-pytorch) + +```bibtex +@inproceedings{reda2022film, + title = {FILM: Frame Interpolation for Large Motion}, + author = {Fitsum Reda and Janne Kontkanen and Eric Tabellion and Deqing Sun and Caroline Pantofaru and Brian Curless}, + booktitle = {European Conference on Computer Vision (ECCV)}, + year = {2022} +} +``` + +```bibtex +@misc{film-tf, + title = {Tensorflow 2 Implementation of "FILM: Frame Interpolation for Large Motion"}, + author = {Fitsum Reda and Janne Kontkanen and Eric Tabellion and Deqing Sun and Caroline Pantofaru and Brian Curless}, + year = {2022}, + publisher = {GitHub}, + journal = {GitHub repository}, + howpublished = {\url{https://github.com/google-research/frame-interpolation}} +} +``` + +### Sepconv +```bibtex +[1] @inproceedings{Niklaus_WACV_2021, + author = {Simon Niklaus and Long Mai and Oliver Wang}, + title = {Revisiting Adaptive Convolutions for Video Frame Interpolation}, + booktitle = {IEEE Winter Conference on Applications of Computer Vision}, + year = {2021} + } +``` + +```bibtex +[2] @inproceedings{Niklaus_ICCV_2017, + author = {Simon Niklaus and Long Mai and Feng Liu}, + title = {Video Frame Interpolation via Adaptive Separable Convolution}, + booktitle = {IEEE International Conference on Computer Vision}, + year = {2017} + } +``` + +```bibtex +[3] @inproceedings{Niklaus_CVPR_2017, + author = {Simon Niklaus and Long Mai and Feng Liu}, + title = {Video Frame Interpolation via Adaptive Convolution}, + booktitle = {IEEE Conference on Computer Vision and Pattern Recognition}, + year = {2017} + } +``` + +### AMT + ```bibtex + @inproceedings{licvpr23amt, + title={AMT: All-Pairs Multi-Field Transforms for Efficient Frame Interpolation}, + author={Li, Zhen and Zhu, Zuo-Liang and Han, Ling-Hao and Hou, Qibin and Guo, Chun-Le and Cheng, Ming-Ming}, + booktitle={IEEE Conference on Computer Vision and Pattern Recognition (CVPR)}, + year={2023} + } + ``` + +### ST-MFNet +```bibtex +@InProceedings{Danier_2022_CVPR, + author = {Danier, Duolikun and Zhang, Fan and Bull, David}, + title = {ST-MFNet: A Spatio-Temporal Multi-Flow Network for Frame Interpolation}, + booktitle = {Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)}, + month = {June}, + year = {2022}, + pages = {3521-3531} +} +``` + +### FLAVR +```bibtex +@article{kalluri2021flavr, + title={FLAVR: Flow-Agnostic Video Representations for Fast Frame Interpolation}, + author={Kalluri, Tarun and Pathak, Deepak and Chandraker, Manmohan and Tran, Du}, + booktitle={arxiv}, + year={2021} +} +``` diff --git a/imports/ComfyUI_Frame_Interpolation/__init__.py b/imports/ComfyUI_Frame_Interpolation/__init__.py new file mode 100644 index 0000000..8e7fb3f --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/__init__.py @@ -0,0 +1,4 @@ +import os +import sys +sys.path.insert(0, os.path.abspath(os.path.dirname(__file__))) + diff --git a/imports/ComfyUI_Frame_Interpolation/config.yaml b/imports/ComfyUI_Frame_Interpolation/config.yaml new file mode 100644 index 0000000..b99d4a7 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/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/imports/ComfyUI_Frame_Interpolation/demo_frames/anime0.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/anime0.png new file mode 100644 index 0000000..14ced4f Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/anime0.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/anime1.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/anime1.png new file mode 100644 index 0000000..1e0c70c Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/anime1.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi0.jpg b/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi0.jpg new file mode 100644 index 0000000..8585b7e Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi0.jpg differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi1.jpg b/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi1.jpg new file mode 100644 index 0000000..27890f7 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/bocchi1.jpg differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/real0.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/real0.png new file mode 100644 index 0000000..6bf29ab Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/real0.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/real1.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/real1.png new file mode 100644 index 0000000..b563e06 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/real1.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00003.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00003.png new file mode 100644 index 0000000..181e260 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00003.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00004.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00004.png new file mode 100644 index 0000000..80ebc6f Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00004.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00005.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00005.png new file mode 100644 index 0000000..a63737d Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/rick/00005.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/violet0.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/violet0.png new file mode 100644 index 0000000..e2aee63 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/violet0.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/demo_frames/violet1.png b/imports/ComfyUI_Frame_Interpolation/demo_frames/violet1.png new file mode 100644 index 0000000..0582b71 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/demo_frames/violet1.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/example.png b/imports/ComfyUI_Frame_Interpolation/example.png new file mode 100644 index 0000000..dbd3701 Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/example.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/import_vfi_utils.py b/imports/ComfyUI_Frame_Interpolation/import_vfi_utils.py new file mode 100644 index 0000000..9c762b0 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/import_vfi_utils.py @@ -0,0 +1,295 @@ +import yaml +import os +from torch.hub import download_url_to_file, get_dir +from urllib.parse import urlparse +import torch +import typing +import traceback +import einops +import gc +import torchvision.transforms.functional as transform +from comfy.model_management import soft_empty_cache, get_torch_device +import numpy as np + +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() + +class InterpolationStateListImport(): + + 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 + + +class MakeInterpolationStateListImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "frame_indices": ("STRING", {"multiline": True, "default": "1,2,3"}), + "is_skip_list": ("BOOLEAN", {"default": True},), + }, + } + + RETURN_TYPES = ("INTERPOLATION_STATES",) + FUNCTION = "create_options" + CATEGORY = "ComfyUI-Frame-Interpolation/VFI" + + def create_options(self, frame_indices: str, is_skip_list: bool): + frame_indices_list = [int(item) for item in frame_indices.split(',')] + + interpolation_state_list = InterpolationStateListImport( + frame_indices=frame_indices_list, + is_skip_list=is_skip_list, + ) + return (interpolation_state_list,) + + +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): + """Load file form http url, will download models if necessary. + + Ref:https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py + + Args: + url (str): URL to be downloaded. + model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir. + Default: None. + progress (bool): Whether to show the download progress. Default: True. + file_name (str): The downloaded file name. If None, use the file name in the url. Default: None. + + Returns: + str: The path to the downloaded file. + """ + 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 + +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 load_file_from_direct_url(model_type, url): + return load_file_from_url(url, get_ckpt_container_path(model_type)) + +def preprocess_frames(frames): + return einops.rearrange(frames[..., :3], "n h w c -> n c h w") + +def postprocess_frames(frames): + return einops.rearrange(frames, "n c h w -> n h w c")[..., :3].cpu() + +def assert_batch_size(frames, batch_size=2, vfi_name=None): + subject_verb = "Most VFI models require" if vfi_name is None else f"VFI model {vfi_name} requires" + assert len(frames) >= batch_size, f"{subject_verb} at least {batch_size} frames to work with, only found {frames.shape[0]}. Please check the frame input using PreviewImage." + +def _generic_frame_loop( + frames, + clear_cache_after_n_frames, + multiplier: typing.Union[typing.SupportsInt, typing.List], + return_middle_frame_function, + *return_middle_frame_function_args, + interpolation_states: InterpolationStateListImport = None, + use_timestep=True, + dtype=torch.float16, + final_logging=True): + + #https://github.com/hzwer/Practical-RIFE/blob/main/inference_video.py#L169 + def non_timestep_inference(frame0, frame1, n): + middle = return_middle_frame_function(frame0, frame1, None, *return_middle_frame_function_args) + if n == 1: + return [middle] + first_half = non_timestep_inference(frame0, middle, n=n//2) + second_half = non_timestep_inference(middle, frame1, n=n//2) + if n%2: + return [*first_half, middle, *second_half] + else: + return [*first_half, *second_half] + + output_frames = torch.zeros(multiplier*frames.shape[0], *frames.shape[1:], dtype=dtype, device="cpu") + out_len = 0 + + number_of_frames_processed_since_last_cleared_cuda_cache = 0 + + for frame_itr in range(len(frames) - 1): # Skip the final frame since there are no frames after it + frame0 = frames[frame_itr:frame_itr+1] + output_frames[out_len] = frame0 # Start with first frame + out_len += 1 + # Ensure that input frames are in fp32 - the same dtype as model + frame0 = frame0.to(dtype=torch.float32) + frame1 = frames[frame_itr+1:frame_itr+2].to(dtype=torch.float32) + + if interpolation_states is not None and interpolation_states.is_frame_skipped(frame_itr): + continue + + # Generate and append a batch of middle frames + middle_frame_batches = [] + + if use_timestep: + for middle_i in range(1, multiplier): + timestep = middle_i/multiplier + + middle_frame = return_middle_frame_function( + frame0.to(DEVICE), + frame1.to(DEVICE), + timestep, + *return_middle_frame_function_args + ).detach().cpu() + middle_frame_batches.append(middle_frame.to(dtype=dtype)) + else: + middle_frames = non_timestep_inference(frame0.to(DEVICE), frame1.to(DEVICE), multiplier - 1) + middle_frame_batches.extend(torch.cat(middle_frames, dim=0).detach().cpu().to(dtype=dtype)) + + # Copy middle frames to output + for middle_frame in middle_frame_batches: + output_frames[out_len] = middle_frame + out_len += 1 + + number_of_frames_processed_since_last_cleared_cuda_cache += 1 + # Try to avoid a memory overflow by clearing cuda cache regularly + if number_of_frames_processed_since_last_cleared_cuda_cache >= clear_cache_after_n_frames: + print("Comfy-VFI: Clearing cache...", end=' ') + soft_empty_cache() + number_of_frames_processed_since_last_cleared_cuda_cache = 0 + print("Done cache clearing") + + gc.collect() + + if final_logging: + print(f"Comfy-VFI done! {len(output_frames)} frames generated at resolution: {output_frames[0].shape}") + # Append final frame + output_frames[out_len] = frames[-1:] + out_len += 1 + # clear cache for courtesy + if final_logging: + print("Comfy-VFI: Final clearing cache...", end = ' ') + soft_empty_cache() + if final_logging: + print("Done cache clearing") + return output_frames[:out_len] + +def generic_frame_loop( + model_name, + frames, + clear_cache_after_n_frames, + multiplier: typing.Union[typing.SupportsInt, typing.List], + return_middle_frame_function, + *return_middle_frame_function_args, + interpolation_states: InterpolationStateListImport = None, + use_timestep=True, + dtype=torch.float32): + + assert_batch_size(frames, vfi_name=model_name.replace('_', ' ').replace('VFI', '')) + if type(multiplier) == int: + return _generic_frame_loop( + frames, + clear_cache_after_n_frames, + multiplier, + return_middle_frame_function, + *return_middle_frame_function_args, + interpolation_states=interpolation_states, + use_timestep=use_timestep, + dtype=dtype + ) + if type(multiplier) == list: + multipliers = list(map(int, multiplier)) + multipliers += [2] * (len(frames) - len(multipliers) - 1) + frame_batches = [] + for frame_itr in range(len(frames) - 1): + multiplier = multipliers[frame_itr] + if multiplier == 0: continue + frame_batch = _generic_frame_loop( + frames[frame_itr:frame_itr+2], + clear_cache_after_n_frames, + multiplier, + return_middle_frame_function, + *return_middle_frame_function_args, + interpolation_states=interpolation_states, + use_timestep=use_timestep, + dtype=dtype, + final_logging=False + ) + if frame_itr != len(frames) - 2: # Not append last frame unless this batch is the last one + frame_batch = frame_batch[:-1] + frame_batches.append(frame_batch) + output_frames = torch.cat(frame_batches) + print(f"Comfy-VFI done! {len(output_frames)} frames generated at resolution: {output_frames[0].shape}") + return output_frames + raise NotImplementedError(f"multipiler of {type(multiplier)}") + +class FloatToIntImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "float": ("FLOAT", {"default": 0, 'min': 0, 'step': 0.01}) + } + } + + RETURN_TYPES = ("INT",) + FUNCTION = "convert" + CATEGORY = "ComfyUI-Frame-Interpolation" + + def convert(self, float): + if hasattr(float, "__iter__"): + return (list(map(int, float)),) + return (int(float),) + +""" def generic_4frame_loop( + frames, + clear_cache_after_n_frames, + multiplier: typing.SupportsInt, + return_middle_frame_function, + *return_middle_frame_function_args, + interpolation_states: InterpolationStateList = None, + use_timestep=False): + + if use_timestep: raise NotImplementedError("Timestep 4 frame VFI model") + def non_timestep_inference(frame_0, frame_1, frame_2, frame_3, n): + middle = return_middle_frame_function(frame_0, frame_1, None, *return_middle_frame_function_args) + if n == 1: + return [middle] + first_half = non_timestep_inference(frame_0, middle, n=n//2) + second_half = non_timestep_inference(middle, frame_1, n=n//2) + if n%2: + return [*first_half, middle, *second_half] + else: + return [*first_half, *second_half] """ \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/install-taichi.bat b/imports/ComfyUI_Frame_Interpolation/install-taichi.bat new file mode 100644 index 0000000..d601f71 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/install-taichi.bat @@ -0,0 +1,11 @@ +@echo off +echo Installing Taichi lang backend... + +if exist "%python_exec%" ( + %python_exec% -s -m pip install taichi +) else ( + echo Installing with system Python + pip install taichi +) + +pause \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/install.bat b/imports/ComfyUI_Frame_Interpolation/install.bat new file mode 100644 index 0000000..84e0f7e --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/install.bat @@ -0,0 +1,16 @@ +@echo off + +set "requirements_txt=%~dp0\requirements-no-cupy.txt" +set "python_exec=..\..\..\python_embeded\python.exe" + +echo Installing ComfyUI Frame Interpolation.. + +if exist "%python_exec%" ( + echo Installing with ComfyUI Portable + %python_exec% -s install.py +) else ( + echo Installing with system Python + python install.py +) + +pause \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/install.py b/imports/ComfyUI_Frame_Interpolation/install.py new file mode 100644 index 0000000..ecbe35d --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/install.py @@ -0,0 +1,59 @@ +import os +from pathlib import Path +import sys +import platform + +def get_cuda_ver_from_dir(cuda_home): + nvrtc = filter(lambda lib_file: "nvrtc-builtins" in lib_file, os.listdir(cuda_home)) + nvrtc = list(nvrtc) + if len(nvrtc) == 0: + return + nvrtc = nvrtc[0] + if ('102' in nvrtc) or ('10.2' in nvrtc): + return '102' + if '110' in nvrtc or ('11.0' in nvrtc): + return '110' + if '111' in nvrtc or ('11.1' in nvrtc): + return '111' + if '11' in nvrtc: + return '11x' + if '12' in nvrtc: + return '12x' + +s_param = '-s' if "python_embeded" in sys.executable else '' + +def get_cuda_home_path(): + if "CUDA_HOME" in os.environ: + return os.environ["CUDA_HOME"] + import torch + torch_lib_path = Path(torch.__file__).parent / "lib" + torch_lib_path = str(torch_lib_path.resolve()) + if os.path.exists(torch_lib_path): + nvrtc = filter(lambda lib_file: "nvrtc-builtins" in lib_file, os.listdir(torch_lib_path)) + nvrtc = list(nvrtc) + return torch_lib_path if len(nvrtc) > 0 else None + +def install_cupy(): + cuda_home = get_cuda_home_path() + try: + if cuda_home is not None: + os.environ["CUDA_HOME"] = cuda_home + os.environ["CUDA_PATH"] = cuda_home + import cupy + print("CuPy is already installed.") + except: + print("Uninstall cupy if existed...") + os.system(f'"{sys.executable}" {s_param} -m pip uninstall -y cupy-wheel cupy-cuda102 cupy-cuda110 cupy-cuda111 cupy-cuda11x cupy-cuda12x') + print("Installing cupy...") + cuda_ver = get_cuda_ver_from_dir(cuda_home) + cupy_package = f"cupy-cuda{cuda_ver}" if cuda_ver is not None else "cupy-wheel" + os.system(f'"{sys.executable}" {s_param} -m pip install {cupy_package}') + +with open(Path(__file__).parent / "requirements-no-cupy.txt", 'r') as f: + for package in f.readlines(): + package = package.strip() + print(f"Installing {package}...") + os.system(f'"{sys.executable}" {s_param} -m pip install {package}') + +print("Checking cupy...") +install_cupy() \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/interpolation_schedule.png b/imports/ComfyUI_Frame_Interpolation/interpolation_schedule.png new file mode 100644 index 0000000..ff92add Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/interpolation_schedule.png differ diff --git a/imports/ComfyUI_Frame_Interpolation/other_nodes.py b/imports/ComfyUI_Frame_Interpolation/other_nodes.py new file mode 100644 index 0000000..75b76fe --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/other_nodes.py @@ -0,0 +1,88 @@ +import latent_preview +import comfy +import einops +import torch + +def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False): + device = comfy.model_management.get_torch_device() + latent_image = latent["samples"] + + if disable_noise: + noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + else: + batch_inds = latent["batch_index"] if "batch_index" in latent else None + noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds) + + noise_mask = None + if "noise_mask" in latent: + noise_mask = latent["noise_mask"] + + preview_format = "JPEG" + if preview_format not in ["JPEG", "PNG"]: + preview_format = "JPEG" + + previewer = latent_preview.get_previewer(device, model.model.latent_format) + + pbar = comfy.utils.ProgressBar(steps) + def callback(step, x0, x, total_steps): + preview_bytes = None + if previewer: + preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0) + pbar.update_absolute(step + 1, total_steps, preview_bytes) + + samples = comfy.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, + denoise=denoise, disable_noise=disable_noise, start_step=start_step, last_step=last_step, + force_full_denoise=force_full_denoise, noise_mask=noise_mask, callback=callback, seed=seed) + out = latent.copy() + out["samples"] = samples + return (out, ) + +class Gradually_More_Denoise_KSampler: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"model": ("MODEL",), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "latent_image": ("LATENT", ), + + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), + "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), + "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), + "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), + + "start_denoise": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "denoise_increment": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.1}), + "denoise_increment_steps": ("INT", {"default": 20, "min": 1, "max": 10000}) + }, + "optional": { "optional_vae": ("VAE",) } + } + + RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "LATENT", "VAE", ) + RETURN_NAMES = ("MODEL", "CONDITIONING+", "CONDITIONING-", "LATENT", "VAE", ) + OUTPUT_NODE = True + FUNCTION = "sample" + CATEGORY = "ComfyUI-Frame-Interpolation/others" + + def sample(self, model, positive, negative, latent_image, optional_vae, + seed, steps, cfg, sampler_name, scheduler,start_denoise, denoise_increment, denoise_increment_steps): + if start_denoise + denoise_increment * denoise_increment_steps > 1.0: + raise Exception(f"Max denoise strength can't over 1.0 (start_denoise={start_denoise}, denoise_increment={denoise_increment}, denoise_increment_steps={denoise_increment_steps}") + + copied_latent = latent_image.copy() + out_samples = [] + + for latent_sample in copied_latent["samples"]: + latent = {"samples": einops.rearrange(latent_sample, "c h w -> 1 c h w")} + #Latent's shape is NCHW + gradually_denoising_samples = [ + common_ksampler( + model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=start_denoise + denoise_increment * i + )[0]["samples"] + for i in range(denoise_increment_steps) + ] + out_samples.extend(gradually_denoising_samples) + + copied_latent["samples"] = torch.cat(out_samples, dim=0) + return (model, positive, negative, copied_latent, optional_vae) \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/requirements-no-cupy.txt b/imports/ComfyUI_Frame_Interpolation/requirements-no-cupy.txt new file mode 100644 index 0000000..c490ca5 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/requirements-no-cupy.txt @@ -0,0 +1,9 @@ +torch +numpy +einops +opencv-contrib-python +kornia +scipy +Pillow +torchvision +tqdm \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/requirements-with-cupy.txt b/imports/ComfyUI_Frame_Interpolation/requirements-with-cupy.txt new file mode 100644 index 0000000..bdfeb47 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/requirements-with-cupy.txt @@ -0,0 +1,10 @@ +torch +numpy +einops +opencv-contrib-python +kornia +scipy +Pillow +torchvision +tqdm +cupy-wheel \ No newline at end of file diff --git a/imports/ComfyUI_Frame_Interpolation/test_vfi_schedule.gif b/imports/ComfyUI_Frame_Interpolation/test_vfi_schedule.gif new file mode 100644 index 0000000..9097cdd Binary files /dev/null and b/imports/ComfyUI_Frame_Interpolation/test_vfi_schedule.gif differ diff --git a/imports/ComfyUI_Frame_Interpolation/vfi_models/film/__init__.py b/imports/ComfyUI_Frame_Interpolation/vfi_models/film/__init__.py new file mode 100644 index 0000000..72b6cea --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/vfi_models/film/__init__.py @@ -0,0 +1,113 @@ +import torch +from comfy.model_management import get_torch_device, soft_empty_cache +import bisect +import numpy as np +import typing +from import_vfi_utils import InterpolationStateListImport, load_file_from_github_release, preprocess_frames, postprocess_frames +import pathlib +import gc + +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] + +class FILM_VFIImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_name": (["film_net_fp32.pt"], ), + "frames": ("IMAGE", ), + "clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}), + "multiplier": ("INT", {"default": 2, "min": 2, "max": 1000}), + }, + "optional": { + "optional_interpolation_states": ("INTERPOLATION_STATES", ) + } + } + + RETURN_TYPES = ("IMAGE", ) + FUNCTION = "vfi" + CATEGORY = "ComfyUI-Frame-Interpolation/VFI" + + def vfi( + self, + ckpt_name: typing.AnyStr, + frames: torch.Tensor, + clear_cache_after_n_frames = 10, + multiplier: typing.SupportsInt = 2, + optional_interpolation_states: InterpolationStateListImport = None, + **kwargs + ): + interpolation_states = optional_interpolation_states + model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name) + model = torch.jit.load(model_path, map_location='cpu') + model.eval() + model = model.to(DEVICE) + dtype = torch.float32 + + frames = preprocess_frames(frames) + number_of_frames_processed_since_last_cleared_cuda_cache = 0 + output_frames = [] + + if type(multiplier) == int: + multipliers = [multiplier] * len(frames) + else: + multipliers = list(map(int, multiplier)) + multipliers += [2] * (len(frames) - len(multipliers) - 1) + for frame_itr in range(len(frames) - 1): # Skip the final frame since there are no frames after it + if interpolation_states is not None and interpolation_states.is_frame_skipped(frame_itr): + continue + #Ensure that input frames are in fp32 - the same dtype as model + frame_0 = frames[frame_itr:frame_itr+1].to(DEVICE).float() + frame_1 = frames[frame_itr+1:frame_itr+2].to(DEVICE).float() + relust = inference(model, frame_0, frame_1, multipliers[frame_itr] - 1) + output_frames.extend([frame.detach().cpu().to(dtype=dtype) for frame in relust[:-1]]) + + number_of_frames_processed_since_last_cleared_cuda_cache += 1 + # Try to avoid a memory overflow by clearing cuda cache regularly + if number_of_frames_processed_since_last_cleared_cuda_cache >= clear_cache_after_n_frames: + print("Comfy-VFI: Clearing cache...", end = ' ') + soft_empty_cache() + number_of_frames_processed_since_last_cleared_cuda_cache = 0 + print("Done cache clearing") + gc.collect() + + output_frames.append(frames[-1:].to(dtype=dtype)) # Append final frame + output_frames = [frame.cpu() for frame in output_frames] #Ensure all frames are in cpu + out = torch.cat(output_frames, dim=0) + # clear cache for courtesy + print("Comfy-VFI: Final clearing cache...", end = ' ') + soft_empty_cache() + print("Done cache clearing") + return (postprocess_frames(out), ) diff --git a/imports/ComfyUI_Frame_Interpolation/vfi_models/film/film_arch.py b/imports/ComfyUI_Frame_Interpolation/vfi_models/film/film_arch.py new file mode 100644 index 0000000..79d0ff5 --- /dev/null +++ b/imports/ComfyUI_Frame_Interpolation/vfi_models/film/film_arch.py @@ -0,0 +1,788 @@ +""" +https://github.com/dajes/frame-interpolation-pytorch/blob/main/feature_extractor.py +https://github.com/dajes/frame-interpolation-pytorch/blob/main/fusion.py +https://github.com/dajes/frame-interpolation-pytorch/blob/main/interpolator.py +https://github.com/dajes/frame-interpolation-pytorch/blob/main/pyramid_flow_estimator.py +https://github.com/dajes/frame-interpolation-pytorch/blob/main/util.py +""" + +"""PyTorch layer for extracting image features for the film_net interpolator. + +The feature extractor implemented here converts an image pyramid into a pyramid +of deep features. The feature pyramid serves a similar purpose as U-Net +architecture's encoder, but we use a special cascaded architecture described in +Multi-view Image Fusion [1]. + +For comprehensiveness, below is a short description of the idea. While the +description is a bit involved, the cascaded feature pyramid can be used just +like any image feature pyramid. + +Why cascaded architeture? +========================= +To understand the concept it is worth reviewing a traditional feature pyramid +first: *A traditional feature pyramid* as in U-net or in many optical flow +networks is built by alternating between convolutions and pooling, starting +from the input image. + +It is well known that early features of such architecture correspond to low +level concepts such as edges in the image whereas later layers extract +semantically higher level concepts such as object classes etc. In other words, +the meaning of the filters in each resolution level is different. For problems +such as semantic segmentation and many others this is a desirable property. + +However, the asymmetric features preclude sharing weights across resolution +levels in the feature extractor itself and in any subsequent neural networks +that follow. This can be a downside, since optical flow prediction, for +instance is symmetric across resolution levels. The cascaded feature +architecture addresses this shortcoming. + +How is it built? +================ +The *cascaded* feature pyramid contains feature vectors that have constant +length and meaning on each resolution level, except few of the finest ones. The +advantage of this is that the subsequent optical flow layer can learn +synergically from many resolutions. This means that coarse level prediction can +benefit from finer resolution training examples, which can be useful with +moderately sized datasets to avoid overfitting. + +The cascaded feature pyramid is built by extracting shallower subtree pyramids, +each one of them similar to the traditional architecture. Each subtree +pyramid S_i is extracted starting from each resolution level: + +image resolution 0 -> S_0 +image resolution 1 -> S_1 +image resolution 2 -> S_2 +... + +If we denote the features at level j of subtree i as S_i_j, the cascaded pyramid +is constructed by concatenating features as follows (assuming subtree depth=3): + +lvl +feat_0 = concat( S_0_0 ) +feat_1 = concat( S_1_0 S_0_1 ) +feat_2 = concat( S_2_0 S_1_1 S_0_2 ) +feat_3 = concat( S_3_0 S_2_1 S_1_2 ) +feat_4 = concat( S_4_0 S_3_1 S_2_2 ) +feat_5 = concat( S_5_0 S_4_1 S_3_2 ) + .... + +In above, all levels except feat_0 and feat_1 have the same number of features +with similar semantic meaning. This enables training a single optical flow +predictor module shared by levels 2,3,4,5... . For more details and evaluation +see [1]. + +[1] Multi-view Image Fusion, Trinidad et al. 2019 +""" +from typing import List + +import torch +from torch import nn +from torch.nn import functional as F + + +class SubTreeExtractorImport(nn.Module): + """Extracts a hierarchical set of features from an image. + + This is a conventional, hierarchical image feature extractor, that extracts + [k, k*2, k*4... ] filters for the image pyramid where k=options.sub_levels. + Each level is followed by average pooling. + """ + + def __init__(self, in_channels=3, channels=64, n_layers=4): + super().__init__() + convs = [] + for i in range(n_layers): + convs.append(nn.Sequential( + conv(in_channels, (channels << i), 3), + conv((channels << i), (channels << i), 3) + )) + in_channels = channels << i + self.convs = nn.ModuleList(convs) + + def forward(self, image: torch.Tensor, n: int) -> List[torch.Tensor]: + """Extracts a pyramid of features from the image. + + Args: + image: TORCH.Tensor with shape BATCH_SIZE x HEIGHT x WIDTH x CHANNELS. + n: number of pyramid levels to extract. This can be less or equal to + options.sub_levels given in the __init__. + Returns: + The pyramid of features, starting from the finest level. Each element + contains the output after the last convolution on the corresponding + pyramid level. + """ + head = image + pyramid = [] + for i, layer in enumerate(self.convs): + head = layer(head) + pyramid.append(head) + if i < n - 1: + head = F.avg_pool2d(head, kernel_size=2, stride=2) + return pyramid + + +class FeatureExtractorImport(nn.Module): + """Extracts features from an image pyramid using a cascaded architecture. + """ + + def __init__(self, in_channels=3, channels=64, sub_levels=4): + super().__init__() + self.extract_sublevels = SubTreeExtractorImport(in_channels, channels, sub_levels) + self.sub_levels = sub_levels + + def forward(self, image_pyramid: List[torch.Tensor]) -> List[torch.Tensor]: + """Extracts a cascaded feature pyramid. + + Args: + image_pyramid: Image pyramid as a list, starting from the finest level. + Returns: + A pyramid of cascaded features. + """ + sub_pyramids: List[List[torch.Tensor]] = [] + for i in range(len(image_pyramid)): + # At each level of the image pyramid, creates a sub_pyramid of features + # with 'sub_levels' pyramid levels, re-using the same SubTreeExtractor. + # We use the same instance since we want to share the weights. + # + # However, we cap the depth of the sub_pyramid so we don't create features + # that are beyond the coarsest level of the cascaded feature pyramid we + # want to generate. + capped_sub_levels = min(len(image_pyramid) - i, self.sub_levels) + sub_pyramids.append(self.extract_sublevels(image_pyramid[i], capped_sub_levels)) + # Below we generate the cascades of features on each level of the feature + # pyramid. Assuming sub_levels=3, The layout of the features will be + # as shown in the example on file documentation above. + feature_pyramid: List[torch.Tensor] = [] + for i in range(len(image_pyramid)): + features = sub_pyramids[i][0] + for j in range(1, self.sub_levels): + if j <= i: + features = torch.cat([features, sub_pyramids[i - j][j]], dim=1) + feature_pyramid.append(features) + return feature_pyramid + + + + + + + + + + + +"""The final fusion stage for the film_net frame interpolator. + +The inputs to this module are the warped input images, image features and +flow fields, all aligned to the target frame (often midway point between the +two original inputs). The output is the final image. FILM has no explicit +occlusion handling -- instead using the abovementioned information this module +automatically decides how to best blend the inputs together to produce content +in areas where the pixels can only be borrowed from one of the inputs. + +Similarly, this module also decides on how much to blend in each input in case +of fractional timestep that is not at the halfway point. For example, if the two +inputs images are at t=0 and t=1, and we were to synthesize a frame at t=0.1, +it often makes most sense to favor the first input. However, this is not +always the case -- in particular in occluded pixels. + +The architecture of the Fusion module follows U-net [1] architecture's decoder +side, e.g. each pyramid level consists of concatenation with upsampled coarser +level output, and two 3x3 convolutions. + +The upsampling is implemented as 'resize convolution', e.g. nearest neighbor +upsampling followed by 2x2 convolution as explained in [2]. The classic U-net +uses max-pooling which has a tendency to create checkerboard artifacts. + +[1] Ronneberger et al. U-Net: Convolutional Networks for Biomedical Image + Segmentation, 2015, https://arxiv.org/pdf/1505.04597.pdf +[2] https://distill.pub/2016/deconv-checkerboard/ +""" +from typing import List + +import torch +from torch import nn +from torch.nn import functional as F + + +_NUMBER_OF_COLOR_CHANNELS = 3 + + +def get_channels_at_level(level, filters): + n_images = 2 + channels = _NUMBER_OF_COLOR_CHANNELS + flows = 2 + + return (sum(filters << i for i in range(level)) + channels + flows) * n_images + + +class FusionImport(nn.Module): + """The decoder.""" + + def __init__(self, n_layers=4, specialized_layers=3, filters=64): + """ + Args: + m: specialized levels + """ + super().__init__() + + # The final convolution that outputs RGB: + self.output_conv = nn.Conv2d(filters, 3, kernel_size=1) + + # Each item 'convs[i]' will contain the list of convolutions to be applied + # for pyramid level 'i'. + self.convs = nn.ModuleList() + + # Create the convolutions. Roughly following the feature extractor, we + # double the number of filters when the resolution halves, but only up to + # the specialized_levels, after which we use the same number of filters on + # all levels. + # + # We create the convs in fine-to-coarse order, so that the array index + # for the convs will correspond to our normal indexing (0=finest level). + # in_channels: tuple = (128, 202, 256, 522, 512, 1162, 1930, 2442) + + in_channels = get_channels_at_level(n_layers, filters) + increase = 0 + for i in range(n_layers)[::-1]: + num_filters = (filters << i) if i < specialized_layers else (filters << specialized_layers) + convs = nn.ModuleList([ + conv(in_channels, num_filters, size=2, activation=None), + conv(in_channels + (increase or num_filters), num_filters, size=3), + conv(num_filters, num_filters, size=3)] + ) + self.convs.append(convs) + in_channels = num_filters + increase = get_channels_at_level(i, filters) - num_filters // 2 + + def forward(self, pyramid: List[torch.Tensor]) -> torch.Tensor: + """Runs the fusion module. + + Args: + pyramid: The input feature pyramid as list of tensors. Each tensor being + in (B x H x W x C) format, with finest level tensor first. + + Returns: + A batch of RGB images. + Raises: + ValueError, if len(pyramid) != config.fusion_pyramid_levels as provided in + the constructor. + """ + + # As a slight difference to a conventional decoder (e.g. U-net), we don't + # apply any extra convolutions to the coarsest level, but just pass it + # to finer levels for concatenation. This choice has not been thoroughly + # evaluated, but is motivated by the educated guess that the fusion part + # probably does not need large spatial context, because at this point the + # features are spatially aligned by the preceding warp. + net = pyramid[-1] + + # Loop starting from the 2nd coarsest level: + # for i in reversed(range(0, len(pyramid) - 1)): + for k, layers in enumerate(self.convs): + i = len(self.convs) - 1 - k + # Resize the tensor from coarser level to match for concatenation. + level_size = pyramid[i].shape[2:4] + net = F.interpolate(net, size=level_size, mode='nearest') + net = layers[0](net) + net = torch.cat([pyramid[i], net], dim=1) + net = layers[1](net) + net = layers[2](net) + net = self.output_conv(net) + return net + + + + + + + + + + + +"""The film_net frame interpolator main model code. + +Basics +====== +The film_net is an end-to-end learned neural frame interpolator implemented as +a PyTorch model. It has the following inputs and outputs: + +Inputs: + x0: image A. + x1: image B. + time: desired sub-frame time. + +Outputs: + image: the predicted in-between image at the chosen time in range [0, 1]. + +Additional outputs include forward and backward warped image pyramids, flow +pyramids, etc., that can be visualized for debugging and analysis. + +Note that many training sets only contain triplets with ground truth at +time=0.5. If a model has been trained with such training set, it will only work +well for synthesizing frames at time=0.5. Such models can only generate more +in-between frames using recursion. + +Architecture +============ +The inference consists of three main stages: 1) feature extraction 2) warping +3) fusion. On high-level, the architecture has similarities to Context-aware +Synthesis for Video Frame Interpolation [1], but the exact architecture is +closer to Multi-view Image Fusion [2] with some modifications for the frame +interpolation use-case. + +Feature extraction stage employs the cascaded multi-scale architecture described +in [2]. The advantage of this architecture is that coarse level flow prediction +can be learned from finer resolution image samples. This is especially useful +to avoid overfitting with moderately sized datasets. + +The warping stage uses a residual flow prediction idea that is similar to +PWC-Net [3], Multi-view Image Fusion [2] and many others. + +The fusion stage is similar to U-Net's decoder where the skip connections are +connected to warped image and feature pyramids. This is described in [2]. + +Implementation Conventions +==================== +Pyramids +-------- +Throughtout the model, all image and feature pyramids are stored as python lists +with finest level first followed by downscaled versions obtained by successively +halving the resolution. The depths of all pyramids are determined by +options.pyramid_levels. The only exception to this is internal to the feature +extractor, where smaller feature pyramids are temporarily constructed with depth +options.sub_levels. + +Color ranges & gamma +-------------------- +The model code makes no assumptions on whether the images are in gamma or +linearized space or what is the range of RGB color values. So a model can be +trained with different choices. This does not mean that all the choices lead to +similar results. In practice the model has been proven to work well with RGB +scale = [0,1] with gamma-space images (i.e. not linearized). + +[1] Context-aware Synthesis for Video Frame Interpolation, Niklaus and Liu, 2018 +[2] Multi-view Image Fusion, Trinidad et al, 2019 +[3] PWC-Net: CNNs for Optical Flow Using Pyramid, Warping, and Cost Volume +""" +from typing import Dict, List + +import torch +from torch import nn + + + +class InterpolatorImport(nn.Module): + def __init__( + self, + pyramid_levels=7, + fusion_pyramid_levels=5, + specialized_levels=3, + sub_levels=4, + filters=64, + flow_convs=(3, 3, 3, 3), + flow_filters=(32, 64, 128, 256), + ): + super().__init__() + self.pyramid_levels = pyramid_levels + self.fusion_pyramid_levels = fusion_pyramid_levels + + self.extract = FeatureExtractorImport(3, filters, sub_levels) + self.predict_flow = PyramidFlowEstimatorImport(filters, flow_convs, flow_filters) + self.fuse = FusionImport(sub_levels, specialized_levels, filters) + + def shuffle_images(self, x0, x1): + return [ + build_image_pyramid(x0, self.pyramid_levels), + build_image_pyramid(x1, self.pyramid_levels) + ] + + def debug_forward(self, x0, x1, batch_dt) -> Dict[str, List[torch.Tensor]]: + image_pyramids = self.shuffle_images(x0, x1) + + # Siamese feature pyramids: + feature_pyramids = [self.extract(image_pyramids[0]), self.extract(image_pyramids[1])] + + # Predict forward flow. + forward_residual_flow_pyramid = self.predict_flow(feature_pyramids[0], feature_pyramids[1]) + + # Predict backward flow. + backward_residual_flow_pyramid = self.predict_flow(feature_pyramids[1], feature_pyramids[0]) + + # Concatenate features and images: + + # Note that we keep up to 'fusion_pyramid_levels' levels as only those + # are used by the fusion module. + + forward_flow_pyramid = flow_pyramid_synthesis(forward_residual_flow_pyramid)[:self.fusion_pyramid_levels] + + backward_flow_pyramid = flow_pyramid_synthesis(backward_residual_flow_pyramid)[:self.fusion_pyramid_levels] + + # We multiply the flows with t and 1-t to warp to the desired fractional time. + # + # Note: In film_net we fix time to be 0.5, and recursively invoke the interpo- + # lator for multi-frame interpolation. Below, we create a constant tensor of + # shape [B]. We use the `time` tensor to infer the batch size. + mid_time = torch.full_like(batch_dt, .5) + backward_flow = multiply_pyramid(backward_flow_pyramid, mid_time[:, 0]) + forward_flow = multiply_pyramid(forward_flow_pyramid, 1 - mid_time[:, 0]) + + pyramids_to_warp = [ + concatenate_pyramids(image_pyramids[0][:self.fusion_pyramid_levels], + feature_pyramids[0][:self.fusion_pyramid_levels]), + concatenate_pyramids(image_pyramids[1][:self.fusion_pyramid_levels], + feature_pyramids[1][:self.fusion_pyramid_levels]) + ] + + # Warp features and images using the flow. Note that we use backward warping + # and backward flow is used to read from image 0 and forward flow from + # image 1. + forward_warped_pyramid = pyramid_warp(pyramids_to_warp[0], backward_flow) + backward_warped_pyramid = pyramid_warp(pyramids_to_warp[1], forward_flow) + + aligned_pyramid = concatenate_pyramids(forward_warped_pyramid, + backward_warped_pyramid) + aligned_pyramid = concatenate_pyramids(aligned_pyramid, backward_flow) + aligned_pyramid = concatenate_pyramids(aligned_pyramid, forward_flow) + + return { + 'image': [self.fuse(aligned_pyramid)], + 'forward_residual_flow_pyramid': forward_residual_flow_pyramid, + 'backward_residual_flow_pyramid': backward_residual_flow_pyramid, + 'forward_flow_pyramid': forward_flow_pyramid, + 'backward_flow_pyramid': backward_flow_pyramid, + } + + + def forward(self, x0, x1, batch_dt) -> torch.Tensor: + return self.debug_forward(x0, x1, batch_dt)['image'][0] + + + + + + + + + + +"""PyTorch layer for estimating optical flow by a residual flow pyramid. + +This approach of estimating optical flow between two images can be traced back +to [1], but is also used by later neural optical flow computation methods such +as SpyNet [2] and PWC-Net [3]. + +The basic idea is that the optical flow is first estimated in a coarse +resolution, then the flow is upsampled to warp the higher resolution image and +then a residual correction is computed and added to the estimated flow. This +process is repeated in a pyramid on coarse to fine order to successively +increase the resolution of both optical flow and the warped image. + +In here, the optical flow predictor is used as an internal component for the +film_net frame interpolator, to warp the two input images into the inbetween, +target frame. + +[1] F. Glazer, Hierarchical motion detection. PhD thesis, 1987. +[2] A. Ranjan and M. J. Black, Optical Flow Estimation using a Spatial Pyramid + Network. 2016 +[3] D. Sun X. Yang, M-Y. Liu and J. Kautz, PWC-Net: CNNs for Optical Flow Using + Pyramid, Warping, and Cost Volume, 2017 +""" +from typing import List + +import torch +from torch import nn +from torch.nn import functional as F + + + +class FlowEstimatorImport(nn.Module): + """Small-receptive field predictor for computing the flow between two images. + + This is used to compute the residual flow fields in PyramidFlowEstimator. + + Note that while the number of 3x3 convolutions & filters to apply is + configurable, two extra 1x1 convolutions are appended to extract the flow in + the end. + + Attributes: + name: The name of the layer + num_convs: Number of 3x3 convolutions to apply + num_filters: Number of filters in each 3x3 convolution + """ + + def __init__(self, in_channels: int, num_convs: int, num_filters: int): + super(FlowEstimatorImport, self).__init__() + + self._convs = nn.ModuleList() + for i in range(num_convs): + self._convs.append(conv(in_channels=in_channels, out_channels=num_filters, size=3)) + in_channels = num_filters + self._convs.append(conv(in_channels, num_filters // 2, size=1)) + in_channels = num_filters // 2 + # For the final convolution, we want no activation at all to predict the + # optical flow vector values. We have done extensive testing on explicitly + # bounding these values using sigmoid, but it turned out that having no + # activation gives better results. + self._convs.append(conv(in_channels, 2, size=1, activation=None)) + + def forward(self, features_a: torch.Tensor, features_b: torch.Tensor) -> torch.Tensor: + """Estimates optical flow between two images. + + Args: + features_a: per pixel feature vectors for image A (B x H x W x C) + features_b: per pixel feature vectors for image B (B x H x W x C) + + Returns: + A tensor with optical flow from A to B + """ + net = torch.cat([features_a, features_b], dim=1) + for conv in self._convs: + net = conv(net) + return net + + +class PyramidFlowEstimatorImport(nn.Module): + """Predicts optical flow by coarse-to-fine refinement. + """ + + def __init__(self, filters: int = 64, + flow_convs: tuple = (3, 3, 3, 3), + flow_filters: tuple = (32, 64, 128, 256)): + super(PyramidFlowEstimatorImport, self).__init__() + + in_channels = filters << 1 + predictors = [] + for i in range(len(flow_convs)): + predictors.append( + FlowEstimatorImport( + in_channels=in_channels, + num_convs=flow_convs[i], + num_filters=flow_filters[i])) + in_channels += filters << (i + 2) + self._predictor = predictors[-1] + self._predictors = nn.ModuleList(predictors[:-1][::-1]) + + def forward(self, feature_pyramid_a: List[torch.Tensor], + feature_pyramid_b: List[torch.Tensor]) -> List[torch.Tensor]: + """Estimates residual flow pyramids between two image pyramids. + + Each image pyramid is represented as a list of tensors in fine-to-coarse + order. Each individual image is represented as a tensor where each pixel is + a vector of image features. + + flow_pyramid_synthesis can be used to convert the residual flow + pyramid returned by this method into a flow pyramid, where each level + encodes the flow instead of a residual correction. + + Args: + feature_pyramid_a: image pyramid as a list in fine-to-coarse order + feature_pyramid_b: image pyramid as a list in fine-to-coarse order + + Returns: + List of flow tensors, in fine-to-coarse order, each level encoding the + difference against the bilinearly upsampled version from the coarser + level. The coarsest flow tensor, e.g. the last element in the array is the + 'DC-term', e.g. not a residual (alternatively you can think of it being a + residual against zero). + """ + levels = len(feature_pyramid_a) + v = self._predictor(feature_pyramid_a[-1], feature_pyramid_b[-1]) + residuals = [v] + for i in range(levels - 2, len(self._predictors) - 1, -1): + # Upsamples the flow to match the current pyramid level. Also, scales the + # magnitude by two to reflect the new size. + level_size = feature_pyramid_a[i].shape[2:4] + v = F.interpolate(2 * v, size=level_size, mode='bilinear') + # Warp feature_pyramid_b[i] image based on the current flow estimate. + warped = warp(feature_pyramid_b[i], v) + # Estimate the residual flow between pyramid_a[i] and warped image: + v_residual = self._predictor(feature_pyramid_a[i], warped) + residuals.insert(0, v_residual) + v = v_residual + v + + for k, predictor in enumerate(self._predictors): + i = len(self._predictors) - 1 - k + # Upsamples the flow to match the current pyramid level. Also, scales the + # magnitude by two to reflect the new size. + level_size = feature_pyramid_a[i].shape[2:4] + v = F.interpolate(2 * v, size=level_size, mode='bilinear') + # Warp feature_pyramid_b[i] image based on the current flow estimate. + warped = warp(feature_pyramid_b[i], v) + # Estimate the residual flow between pyramid_a[i] and warped image: + v_residual = predictor(feature_pyramid_a[i], warped) + residuals.insert(0, v_residual) + v = v_residual + v + return residuals + + + + + + + + + + +"""Various utilities used in the film_net frame interpolator model.""" +from typing import List, Optional + +import cv2 +import numpy as np +import torch +from torch import nn +from torch.nn import functional as F + + +def pad_batch(batch, align): + height, width = batch.shape[1:3] + height_to_pad = (align - height % align) if height % align != 0 else 0 + width_to_pad = (align - width % align) if width % align != 0 else 0 + + crop_region = [height_to_pad >> 1, width_to_pad >> 1, height + (height_to_pad >> 1), width + (width_to_pad >> 1)] + batch = np.pad(batch, ((0, 0), (height_to_pad >> 1, height_to_pad - (height_to_pad >> 1)), + (width_to_pad >> 1, width_to_pad - (width_to_pad >> 1)), (0, 0)), mode='constant') + return batch, crop_region + + +def load_image(path, align=64): + image = cv2.cvtColor(cv2.imread(path), cv2.COLOR_BGR2RGB).astype(np.float32) / np.float32(255) + image_batch, crop_region = pad_batch(np.expand_dims(image, axis=0), align) + return image_batch, crop_region + + +def build_image_pyramid(image: torch.Tensor, pyramid_levels: int = 3) -> List[torch.Tensor]: + """Builds an image pyramid from a given image. + + The original image is included in the pyramid and the rest are generated by + successively halving the resolution. + + Args: + image: the input image. + options: film_net options object + + Returns: + A list of images starting from the finest with options.pyramid_levels items + """ + + pyramid = [] + for i in range(pyramid_levels): + pyramid.append(image) + if i < pyramid_levels - 1: + image = F.avg_pool2d(image, 2, 2) + return pyramid + + +def warp(image: torch.Tensor, flow: torch.Tensor) -> torch.Tensor: + """Backward warps the image using the given flow. + + Specifically, the output pixel in batch b, at position x, y will be computed + as follows: + (flowed_y, flowed_x) = (y+flow[b, y, x, 1], x+flow[b, y, x, 0]) + output[b, y, x] = bilinear_lookup(image, b, flowed_y, flowed_x) + + Note that the flow vectors are expected as [x, y], e.g. x in position 0 and + y in position 1. + + Args: + image: An image with shape BxHxWxC. + flow: A flow with shape BxHxWx2, with the two channels denoting the relative + offset in order: (dx, dy). + Returns: + A warped image. + """ + flow = -flow.flip(1) + + dtype = flow.dtype + device = flow.device + + # warped = tfa_image.dense_image_warp(image, flow) + # Same as above but with pytorch + ls1 = 1 - 1 / flow.shape[3] + ls2 = 1 - 1 / flow.shape[2] + + normalized_flow2 = flow.permute(0, 2, 3, 1) / torch.tensor( + [flow.shape[2] * .5, flow.shape[3] * .5], dtype=dtype, device=device)[None, None, None] + normalized_flow2 = torch.stack([ + torch.linspace(-ls1, ls1, flow.shape[3], dtype=dtype, device=device)[None, None, :] - normalized_flow2[..., 1], + torch.linspace(-ls2, ls2, flow.shape[2], dtype=dtype, device=device)[None, :, None] - normalized_flow2[..., 0], + ], dim=3) + + warped = F.grid_sample(image, normalized_flow2, + mode='bilinear', padding_mode='border', align_corners=False) + return warped.reshape(image.shape) + + +def multiply_pyramid(pyramid: List[torch.Tensor], + scalar: torch.Tensor) -> List[torch.Tensor]: + """Multiplies all image batches in the pyramid by a batch of scalars. + + Args: + pyramid: Pyramid of image batches. + scalar: Batch of scalars. + + Returns: + An image pyramid with all images multiplied by the scalar. + """ + # To multiply each image with its corresponding scalar, we first transpose + # the batch of images from BxHxWxC-format to CxHxWxB. This can then be + # multiplied with a batch of scalars, then we transpose back to the standard + # BxHxWxC form. + return [image * scalar for image in pyramid] + + +def flow_pyramid_synthesis( + residual_pyramid: List[torch.Tensor]) -> List[torch.Tensor]: + """Converts a residual flow pyramid into a flow pyramid.""" + flow = residual_pyramid[-1] + flow_pyramid: List[torch.Tensor] = [flow] + for residual_flow in residual_pyramid[:-1][::-1]: + level_size = residual_flow.shape[2:4] + flow = F.interpolate(2 * flow, size=level_size, mode='bilinear') + flow = residual_flow + flow + flow_pyramid.insert(0, flow) + return flow_pyramid + + +def pyramid_warp(feature_pyramid: List[torch.Tensor], + flow_pyramid: List[torch.Tensor]) -> List[torch.Tensor]: + """Warps the feature pyramid using the flow pyramid. + + Args: + feature_pyramid: feature pyramid starting from the finest level. + flow_pyramid: flow fields, starting from the finest level. + + Returns: + Reverse warped feature pyramid. + """ + warped_feature_pyramid = [] + for features, flow in zip(feature_pyramid, flow_pyramid): + warped_feature_pyramid.append(warp(features, flow)) + return warped_feature_pyramid + + +def concatenate_pyramids(pyramid1: List[torch.Tensor], + pyramid2: List[torch.Tensor]) -> List[torch.Tensor]: + """Concatenates each pyramid level together in the channel dimension.""" + result = [] + for features1, features2 in zip(pyramid1, pyramid2): + result.append(torch.cat([features1, features2], dim=1)) + return result + + +def conv(in_channels, out_channels, size, activation: Optional[str] = 'relu'): + # Since PyTorch doesn't have an in-built activation in Conv2d, we use a + # Sequential layer to combine Conv2d and Leaky ReLU in one module. + _conv = nn.Conv2d( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=size, + padding='same') + if activation is None: + return _conv + assert activation == 'relu' + return nn.Sequential( + _conv, + nn.LeakyReLU(.2) + ) \ No newline at end of file