Merge pull request #75 from banodoco/feature/interpolate
Feature/interpolate
@@ -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
|
||||
|
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
|
||||
|
||||

|
||||

|
||||
|
||||
## 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.
|
||||

|
||||
|
||||
### Complex workflow
|
||||
It's used in AnimationDiff (can load workflow metadata)
|
||||

|
||||
|
||||
## 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"
|
||||
|
After Width: | Height: | Size: 333 KiB |
|
After Width: | Height: | Size: 322 KiB |
|
After Width: | Height: | Size: 127 KiB |
|
After Width: | Height: | Size: 136 KiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 446 KiB |
|
After Width: | Height: | Size: 347 KiB |
|
After Width: | Height: | Size: 349 KiB |
|
After Width: | Height: | Size: 868 KiB |
|
After Width: | Height: | Size: 929 KiB |
|
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()
|
||||
|
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
|
||||
|
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)
|
||||
)
|
||||