Merge pull request #75 from banodoco/feature/interpolate

Feature/interpolate
This commit is contained in:
POM
2024-05-25 23:43:55 +02:00
committed by GitHub
30 changed files with 1644 additions and 16 deletions
+28 -14
View File
@@ -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,25 +583,33 @@ 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
# 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 index < images.shape[0]:
images = torch.cat((images[:index], images[index+1:]))
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 🎞️🅢🅜",
}
@@ -0,0 +1,3 @@
ckpts
__pycache__
test_result
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.4 MiB

@@ -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.
@@ -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}
}
```
@@ -0,0 +1,4 @@
import os
import sys
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
@@ -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"
Binary file not shown.

After

Width:  |  Height:  |  Size: 333 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 322 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 127 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 136 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 446 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 347 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 349 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 868 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 929 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 178 KiB

@@ -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] """
@@ -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
@@ -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
@@ -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()
Binary file not shown.

After

Width:  |  Height:  |  Size: 369 KiB

@@ -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)
@@ -0,0 +1,9 @@
torch
numpy
einops
opencv-contrib-python
kornia
scipy
Pillow
torchvision
tqdm
@@ -0,0 +1,10 @@
torch
numpy
einops
opencv-contrib-python
kornia
scipy
Pillow
torchvision
tqdm
cupy-wheel
Binary file not shown.

After

Width:  |  Height:  |  Size: 8.0 MiB

@@ -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), )
@@ -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)
)