FIxing classes

This commit is contained in:
peteromallet
2024-05-25 20:05:19 +02:00
parent d756dcf1e1
commit 8370a2f1ef
54 changed files with 12 additions and 20004 deletions
+3 -2
View File
@@ -10,6 +10,7 @@ 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
@@ -565,7 +566,7 @@ class BatchCreativeInterpolationNode:
# import the class FILM_VFI from ComfyUI-Frame-Interpolation/vfi_models/film/__init__.py
from .imports.ComfyUI_Frame_Interpolation.vfi_models.film import FILM_VFI
# from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport
@@ -592,7 +593,7 @@ class DropFramesByIndex:
frames_to_drop = sorted(frames_to_drop, reverse=True)
# Create instance of FILM_VFI within the function
film_vfi = FILM_VFI() # Assuming FILM_VFI does not require any special setup
film_vfi = FILM_VFIImport() # Assuming FILM_VFI does not require any special setup
for index in frames_to_drop:
if 0 < index < images.shape[0] - 1:
@@ -2,41 +2,3 @@ import os
import sys
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from .other_nodes import Gradually_More_Denoise_KSampler
#Some models are commented out because the code is not completed
#from vfi_models.eisai import EISAI_VFI
from vfi_models.gmfss_fortuna import GMFSS_Fortuna_VFI
from vfi_models.ifrnet import IFRNet_VFI
from vfi_models.ifunet import IFUnet_VFI
from vfi_models.m2m import M2M_VFI
from vfi_models.rife import RIFE_VFI
from vfi_models.sepconv import SepconvVFI
from vfi_models.amt import AMT_VFI
from vfi_models.film import FILM_VFI
from vfi_models.stmfnet import STMFNet_VFI
from vfi_models.flavr import FLAVR_VFI
from vfi_models.cain import CAIN_VFI
from vfi_utils import MakeInterpolationStateList, FloatToInt
NODE_CLASS_MAPPINGS = {
"KSampler Gradually Adding More Denoise (efficient)": Gradually_More_Denoise_KSampler,
# "EISAI VFI": EISAI_VFI,
"GMFSS Fortuna VFI": GMFSS_Fortuna_VFI,
"IFRNet VFI": IFRNet_VFI,
"IFUnet VFI": IFUnet_VFI,
"M2M VFI": M2M_VFI,
"RIFE VFI": RIFE_VFI,
"Sepconv VFI": SepconvVFI,
"AMT VFI": AMT_VFI,
"FILM VFI": FILM_VFI,
"Make Interpolation State List": MakeInterpolationStateList,
"STMFNet VFI": STMFNet_VFI,
"FLAVR VFI": FLAVR_VFI,
"CAIN VFI": CAIN_VFI,
"VFI FloatToInt": FloatToInt
}
NODE_DISPLAY_NAME_MAPPINGS = {
"RIFE VFI": "RIFE VFI (recommend rife47 and rife49)"
}
@@ -1,87 +0,0 @@
import pathlib
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_direct_url, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
from comfy.model_management import get_torch_device
from .amt_arch import AMT_S, AMT_L, AMT_G, InputPadder
#https://github.com/MCG-NKU/AMT/tree/main/cfgs
CKPT_CONFIGS = {
"amt-s.pth": {
"network": AMT_S,
"params": { "corr_radius": 3, "corr_lvls": 4, "num_flows": 3 }
},
"amt-l.pth": {
"network": AMT_L,
"params": { "corr_radius": 3, "corr_lvls": 4, "num_flows": 5 }
},
"amt-g.pth": {
"network": AMT_G,
"params": { "corr_radius": 3, "corr_lvls": 4, "num_flows": 5 }
},
"gopro_amt-s.pth": {
"network": AMT_S,
"params": { "corr_radius": 3, "corr_lvls": 4, "num_flows": 3 }
}
}
MODEL_TYPE = pathlib.Path(__file__).parent.name
class AMT_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (list(CKPT_CONFIGS.keys()), ),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 1, "min": 1, "max": 100}),
"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: typing.SupportsInt = 1,
multiplier: typing.SupportsInt = 2,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
model_path = load_file_from_direct_url(MODEL_TYPE, f"https://huggingface.co/lalala125/AMT/resolve/main/{ckpt_name}")
ckpt_config = CKPT_CONFIGS[ckpt_name]
interpolation_model = ckpt_config["network"](**ckpt_config["params"])
interpolation_model.load_state_dict(torch.load(model_path)["state_dict"])
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
padder = InputPadder(frames.shape, 16)
frames = padder.pad(frames)
def return_middle_frame(frame_0, frame_1, timestep, model):
return model(
frame_0,
frame_1,
embt=torch.FloatTensor([timestep] * frame_0.shape[0]).view(frame_0.shape[0], 1, 1, 1).to(get_torch_device()),
scale_factor=1.0,
eval=True
)["imgt_pred"]
args = [interpolation_model]
out = generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
out = padder.unpad(out)
out = postprocess_frames(out)
return (out,)
File diff suppressed because it is too large Load Diff
@@ -1,64 +0,0 @@
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
from comfy.model_management import get_torch_device
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["pretrained_cain.pth"]
class CAIN_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"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: typing.SupportsInt = 1,
multiplier: typing.SupportsInt = 2,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
from .cain_arch import CAIN
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
sd = torch.load(model_path)["state_dict"]
sd = {key.replace('module.', ''): value for key, value in sd.items()}
global interpolation_model
interpolation_model = CAIN(depth=3)
interpolation_model.load_state_dict(sd)
interpolation_model.eval().to(get_torch_device())
del sd
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model):
#CAIN does some direct modifications to input frame tensors so we need to clone them
return model(frame_0.detach().clone(), frame_1.detach().clone())[0]
args = [interpolation_model]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, use_timestep=False, dtype=torch.float32)
)
return (out,)
@@ -1,74 +0,0 @@
import math
import numpy as np
import torch
import torch.nn as nn
from .common import *
class Encoder(nn.Module):
def __init__(self, in_channels=3, depth=3):
super(Encoder, self).__init__()
# Shuffle pixels to expand in channel dimension
# shuffler_list = [PixelShuffle(0.5) for i in range(depth)]
# self.shuffler = nn.Sequential(*shuffler_list)
self.shuffler = PixelShuffle(1 / 2**depth)
relu = nn.LeakyReLU(0.2, True)
# FF_RCAN or FF_Resblocks
self.interpolate = Interpolation(5, 12, in_channels * (4**depth), act=relu)
def forward(self, x1, x2):
"""
Encoder: Shuffle-spread --> Feature Fusion --> Return fused features
"""
feats1 = self.shuffler(x1)
feats2 = self.shuffler(x2)
feats = self.interpolate(feats1, feats2)
return feats
class Decoder(nn.Module):
def __init__(self, depth=3):
super(Decoder, self).__init__()
# shuffler_list = [PixelShuffle(2) for i in range(depth)]
# self.shuffler = nn.Sequential(*shuffler_list)
self.shuffler = PixelShuffle(2**depth)
def forward(self, feats):
out = self.shuffler(feats)
return out
class CAIN(nn.Module):
def __init__(self, depth=3):
super(CAIN, self).__init__()
self.encoder = Encoder(in_channels=3, depth=depth)
self.decoder = Decoder(depth=depth)
def forward(self, x1, x2):
x1, m1 = sub_mean(x1)
x2, m2 = sub_mean(x2)
if not self.training:
paddingInput, paddingOutput = InOutPaddings(x1)
x1 = paddingInput(x1)
x2 = paddingInput(x2)
feats = self.encoder(x1, x2)
out = self.decoder(feats)
if not self.training:
out = paddingOutput(out)
mi = (m1 + m2) / 2
out += mi
return out, feats
@@ -1,95 +0,0 @@
import math
import numpy as np
import torch
import torch.nn as nn
from .common import *
from comfy.model_management import get_torch_device
class Encoder(nn.Module):
def __init__(self, in_channels=3, depth=3, nf_start=32, norm=False):
super(Encoder, self).__init__()
self.device = get_torch_device()
nf = nf_start
relu = nn.LeakyReLU(negative_slope=0.2, inplace=True)
self.body = nn.Sequential(
ConvNorm(in_channels, nf * 1, 7, stride=1, norm=norm),
relu,
ConvNorm(nf * 1, nf * 2, 5, stride=2, norm=norm),
relu,
ConvNorm(nf * 2, nf * 4, 5, stride=2, norm=norm),
relu,
ConvNorm(nf * 4, nf * 6, 5, stride=2, norm=norm)
)
self.interpolate = Interpolation(5, 12, nf * 6, reduction=16, act=relu)
def forward(self, x1, x2):
"""
Encoder: Feature Extraction --> Feature Fusion --> Return
"""
feats1 = self.body(x1)
feats2 = self.body(x2)
feats = self.interpolate(feats1, feats2)
return feats
class Decoder(nn.Module):
def __init__(self, in_channels=192, out_channels=3, depth=3, norm=False, up_mode='shuffle'):
super(Decoder, self).__init__()
self.device = get_torch_device()
relu = nn.LeakyReLU(negative_slope=0.2, inplace=True)
nf = [in_channels, (in_channels*2)//3, in_channels//3, in_channels//6]
#nf = [192, 128, 64, 32]
#nf = [186, 124, 62, 31]
self.body = nn.Sequential(
UpConvNorm(nf[0], nf[1], mode=up_mode, norm=norm),
ResBlock(nf[1], nf[1], norm=norm, act=relu),
UpConvNorm(nf[1], nf[2], mode=up_mode, norm=norm),
ResBlock(nf[2], nf[2], norm=norm, act=relu),
UpConvNorm(nf[2], nf[3], mode=up_mode, norm=norm),
ResBlock(nf[3], nf[3], norm=norm, act=relu),
conv7x7(nf[3], out_channels)
)
def forward(self, feats):
out = self.body(feats)
#out = self.conv_final(out)
return out
class CAIN_EncDec(nn.Module):
def __init__(self, depth=3, n_resblocks=3, start_filts=32, up_mode='shuffle'):
super(CAIN_EncDec, self).__init__()
self.depth = depth
self.encoder = Encoder(in_channels=3, depth=depth, norm=False)
self.decoder = Decoder(in_channels=start_filts*6, depth=depth, norm=False, up_mode=up_mode)
def forward(self, x1, x2):
x1, m1 = sub_mean(x1)
x2, m2 = sub_mean(x2)
if not self.training:
paddingInput, paddingOutput = InOutPaddings(x1)
x1 = paddingInput(x1)
x2 = paddingInput(x2)
feats = self.encoder(x1, x2)
out = self.decoder(feats)
if not self.training:
out = paddingOutput(out)
mi = (m1 + m2)/2
out += mi
return out, feats
@@ -1,73 +0,0 @@
import math
import numpy as np
import torch
import torch.nn as nn
from .common import *
from comfy.model_management import get_torch_device
class Encoder(nn.Module):
def __init__(self, in_channels=3, depth=3):
super(Encoder, self).__init__()
self.device = get_torch_device()
self.shuffler = PixelShuffle(1/2**depth)
# self.shuffler = nn.Sequential(
# PixelShuffle(1/2),
# PixelShuffle(1/2),
# PixelShuffle(1/2))
self.interpolate = Interpolation_res(5, 12, in_channels * (4**depth))
def forward(self, x1, x2):
feats1 = self.shuffler(x1)
feats2 = self.shuffler(x2)
feats = self.interpolate(feats1, feats2)
return feats
class Decoder(nn.Module):
def __init__(self, depth=3):
super(Decoder, self).__init__()
self.device = get_torch_device()
self.shuffler = PixelShuffle(2**depth)
# self.shuffler = nn.Sequential(
# PixelShuffle(2),
# PixelShuffle(2),
# PixelShuffle(2))
def forward(self, feats):
out = self.shuffler(feats)
return out
class CAIN_NoCA(nn.Module):
def __init__(self, depth=3):
super(CAIN_NoCA, self).__init__()
self.depth = depth
self.encoder = Encoder(in_channels=3, depth=depth)
self.decoder = Decoder(depth=depth)
def forward(self, x1, x2):
x1, m1 = sub_mean(x1)
x2, m2 = sub_mean(x2)
if not self.training:
paddingInput, paddingOutput = InOutPaddings(x1)
x1 = paddingInput(x1)
x2 = paddingInput(x2)
feats = self.encoder(x1, x2)
out = self.decoder(feats)
if not self.training:
out = paddingOutput(out)
mi = (m1 + m2) / 2
out += mi
return out, feats
@@ -1,361 +0,0 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
def sub_mean(x):
mean = x.mean(2, keepdim=True).mean(3, keepdim=True)
x -= mean
return x, mean
def InOutPaddings(x):
w, h = x.size(3), x.size(2)
padding_width, padding_height = 0, 0
if w != ((w >> 7) << 7):
padding_width = (((w >> 7) + 1) << 7) - w
if h != ((h >> 7) << 7):
padding_height = (((h >> 7) + 1) << 7) - h
paddingInput = nn.ReflectionPad2d(padding=[padding_width // 2, padding_width - padding_width // 2,
padding_height // 2, padding_height - padding_height // 2])
paddingOutput = nn.ReflectionPad2d(padding=[0 - padding_width // 2, padding_width // 2 - padding_width,
0 - padding_height // 2, padding_height // 2 - padding_height])
return paddingInput, paddingOutput
class ConvNorm(nn.Module):
def __init__(self, in_feat, out_feat, kernel_size, stride=1, norm=False):
super(ConvNorm, self).__init__()
reflection_padding = kernel_size // 2
self.reflection_pad = nn.ReflectionPad2d(reflection_padding)
self.conv = nn.Conv2d(in_feat, out_feat, stride=stride, kernel_size=kernel_size, bias=True)
self.norm = norm
if norm == 'IN':
self.norm = nn.InstanceNorm2d(out_feat, track_running_stats=True)
elif norm == 'BN':
self.norm = nn.BatchNorm2d(out_feat)
def forward(self, x):
out = self.reflection_pad(x)
out = self.conv(out)
if self.norm:
out = self.norm(out)
return out
class UpConvNorm(nn.Module):
def __init__(self, in_channels, out_channels, mode='transpose', norm=False):
super(UpConvNorm, self).__init__()
if mode == 'transpose':
self.upconv = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=4, stride=2, padding=1)
elif mode == 'shuffle':
self.upconv = nn.Sequential(
ConvNorm(in_channels, 4*out_channels, kernel_size=3, stride=1, norm=norm),
PixelShuffle(2))
else:
# out_channels is always going to be the same as in_channels
self.upconv = nn.Sequential(
nn.Upsample(mode='bilinear', scale_factor=2, align_corners=False),
ConvNorm(in_channels, out_channels, kernel_size=1, stride=1, norm=norm))
def forward(self, x):
out = self.upconv(x)
return out
class meanShift(nn.Module):
def __init__(self, rgbRange, rgbMean, sign, nChannel=3):
super(meanShift, self).__init__()
if nChannel == 1:
l = rgbMean[0] * rgbRange * float(sign)
self.shifter = nn.Conv2d(1, 1, kernel_size=1, stride=1, padding=0)
self.shifter.weight.data = torch.eye(1).view(1, 1, 1, 1)
self.shifter.bias.data = torch.Tensor([l])
elif nChannel == 3:
r = rgbMean[0] * rgbRange * float(sign)
g = rgbMean[1] * rgbRange * float(sign)
b = rgbMean[2] * rgbRange * float(sign)
self.shifter = nn.Conv2d(3, 3, kernel_size=1, stride=1, padding=0)
self.shifter.weight.data = torch.eye(3).view(3, 3, 1, 1)
self.shifter.bias.data = torch.Tensor([r, g, b])
else:
r = rgbMean[0] * rgbRange * float(sign)
g = rgbMean[1] * rgbRange * float(sign)
b = rgbMean[2] * rgbRange * float(sign)
self.shifter = nn.Conv2d(6, 6, kernel_size=1, stride=1, padding=0)
self.shifter.weight.data = torch.eye(6).view(6, 6, 1, 1)
self.shifter.bias.data = torch.Tensor([r, g, b, r, g, b])
# Freeze the meanShift layer
for params in self.shifter.parameters():
params.requires_grad = False
def forward(self, x):
x = self.shifter(x)
return x
""" CONV - (BN) - RELU - CONV - (BN) """
class ResBlock(nn.Module):
def __init__(self, in_feat, out_feat, kernel_size=3, reduction=False, bias=True, # 'reduction' is just for placeholder
norm=False, act=nn.ReLU(True), downscale=False):
super(ResBlock, self).__init__()
self.body = nn.Sequential(
ConvNorm(in_feat, out_feat, kernel_size=kernel_size, stride=2 if downscale else 1),
act,
ConvNorm(out_feat, out_feat, kernel_size=kernel_size, stride=1)
)
self.downscale = None
if downscale:
self.downscale = nn.Conv2d(in_feat, out_feat, kernel_size=1, stride=2)
def forward(self, x):
res = x
out = self.body(x)
if self.downscale is not None:
res = self.downscale(res)
out += res
return out
## Channel Attention (CA) Layer
class CALayer(nn.Module):
def __init__(self, channel, reduction=16):
super(CALayer, self).__init__()
# global average pooling: feature --> point
self.avg_pool = nn.AdaptiveAvgPool2d(1)
# feature channel downscale and upscale --> channel weight
self.conv_du = nn.Sequential(
nn.Conv2d(channel, channel // reduction, 1, padding=0, bias=True),
nn.ReLU(inplace=True),
nn.Conv2d(channel // reduction, channel, 1, padding=0, bias=True),
nn.Sigmoid()
)
def forward(self, x):
y = self.avg_pool(x)
y = self.conv_du(y)
return x * y, y
## Residual Channel Attention Block (RCAB)
class RCAB(nn.Module):
def __init__(self, in_feat, out_feat, kernel_size, reduction, bias=True,
norm=False, act=nn.ReLU(True), downscale=False, return_ca=False):
super(RCAB, self).__init__()
self.body = nn.Sequential(
ConvNorm(in_feat, out_feat, kernel_size, stride=2 if downscale else 1, norm=norm),
act,
ConvNorm(out_feat, out_feat, kernel_size, stride=1, norm=norm),
CALayer(out_feat, reduction)
)
self.downscale = downscale
if downscale:
self.downConv = nn.Conv2d(in_feat, out_feat, kernel_size=3, stride=2, padding=1)
self.return_ca = return_ca
def forward(self, x):
res = x
out, ca = self.body(x)
if self.downscale:
res = self.downConv(res)
out += res
if self.return_ca:
return out, ca
else:
return out
## Residual Group (RG)
class ResidualGroup(nn.Module):
def __init__(self, Block, n_resblocks, n_feat, kernel_size, reduction, act, norm=False):
super(ResidualGroup, self).__init__()
modules_body = [Block(n_feat, n_feat, kernel_size, reduction, bias=True, norm=norm, act=act)
for _ in range(n_resblocks)]
modules_body.append(ConvNorm(n_feat, n_feat, kernel_size, stride=1, norm=norm))
self.body = nn.Sequential(*modules_body)
def forward(self, x):
res = self.body(x)
res += x
return res
def pixel_shuffle(input, scale_factor):
batch_size, channels, in_height, in_width = input.size()
out_channels = int(int(channels / scale_factor) / scale_factor)
out_height = int(in_height * scale_factor)
out_width = int(in_width * scale_factor)
if scale_factor >= 1:
input_view = input.contiguous().view(batch_size, out_channels, scale_factor, scale_factor, in_height, in_width)
shuffle_out = input_view.permute(0, 1, 4, 2, 5, 3).contiguous()
else:
block_size = int(1 / scale_factor)
input_view = input.contiguous().view(batch_size, channels, out_height, block_size, out_width, block_size)
shuffle_out = input_view.permute(0, 1, 3, 5, 2, 4).contiguous()
return shuffle_out.view(batch_size, out_channels, out_height, out_width)
class PixelShuffle(nn.Module):
def __init__(self, scale_factor):
super(PixelShuffle, self).__init__()
self.scale_factor = scale_factor
def forward(self, x):
return pixel_shuffle(x, self.scale_factor)
def extra_repr(self):
return 'scale_factor={}'.format(self.scale_factor)
def conv(in_channels, out_channels, kernel_size,
stride=1, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=kernel_size,
padding=kernel_size//2,
stride=1,
bias=bias,
groups=groups)
def conv1x1(in_channels, out_channels, stride=1, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=1,
stride=stride,
bias=bias,
groups=groups)
def conv3x3(in_channels, out_channels, stride=1,
padding=1, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=3,
stride=stride,
padding=padding,
bias=bias,
groups=groups)
def conv5x5(in_channels, out_channels, stride=1,
padding=2, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=5,
stride=stride,
padding=padding,
bias=bias,
groups=groups)
def conv7x7(in_channels, out_channels, stride=1,
padding=3, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=7,
stride=stride,
padding=padding,
bias=bias,
groups=groups)
def upconv2x2(in_channels, out_channels, mode='shuffle'):
if mode == 'transpose':
return nn.ConvTranspose2d(
in_channels,
out_channels,
kernel_size=4,
stride=2,
padding=1)
elif mode == 'shuffle':
return nn.Sequential(
conv3x3(in_channels, 4*out_channels),
PixelShuffle(2))
else:
# out_channels is always going to be the same as in_channels
return nn.Sequential(
nn.Upsample(mode='bilinear', scale_factor=2, align_corners=False),
conv1x1(in_channels, out_channels))
class Interpolation(nn.Module):
def __init__(self, n_resgroups, n_resblocks, n_feats,
reduction=16, act=nn.LeakyReLU(0.2, True), norm=False):
super(Interpolation, self).__init__()
# define modules: head, body, tail
self.headConv = conv3x3(n_feats * 2, n_feats)
modules_body = [
ResidualGroup(
RCAB,
n_resblocks=n_resblocks,
n_feat=n_feats,
kernel_size=3,
reduction=reduction,
act=act,
norm=norm)
for _ in range(n_resgroups)]
self.body = nn.Sequential(*modules_body)
self.tailConv = conv3x3(n_feats, n_feats)
def forward(self, x0, x1):
# Build input tensor
x = torch.cat([x0, x1], dim=1)
x = self.headConv(x)
res = self.body(x)
res += x
out = self.tailConv(res)
return out
class Interpolation_res(nn.Module):
def __init__(self, n_resgroups, n_resblocks, n_feats,
act=nn.LeakyReLU(0.2, True), norm=False):
super(Interpolation_res, self).__init__()
# define modules: head, body, tail (reduces concatenated inputs to n_feat)
self.headConv = conv3x3(n_feats * 2, n_feats)
modules_body = [ResidualGroup(ResBlock, n_resblocks=n_resblocks, n_feat=n_feats, kernel_size=3,
reduction=0, act=act, norm=norm)
for _ in range(n_resgroups)]
self.body = nn.Sequential(*modules_body)
self.tailConv = conv3x3(n_feats, n_feats)
def forward(self, x0, x1):
# Build input tensor
x = torch.cat([x0, x1], dim=1)
x = self.headConv(x)
res = x
for m in self.body:
res = m(res)
res += x
x = self.tailConv(res)
return x
@@ -1,84 +0,0 @@
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
import torch
import torch.nn as nn
from comfy.model_management import soft_empty_cache, get_torch_device
MODEL_TYPE = pathlib.Path(__file__).parent.name
MODEL_FILE_NAMES = {
"ssl": "eisai_ssl.pt",
"dtm": "eisai_dtm.pt",
"raft": "eisai_anime_interp_full.ckpt"
}
class EISAI(nn.Module):
def __init__(self, model_file_names) -> None:
from .eisai_arch import SoftsplatLite, DTM, RAFT
super(EISAI, self).__init__()
self.raft = RAFT(load_file_from_github_release(MODEL_TYPE, model_file_names["raft"]))
self.raft.to(get_torch_device()).eval()
self.ssl = SoftsplatLite()
self.ssl.load_state_dict(torch.load(load_file_from_github_release(MODEL_TYPE, model_file_names["ssl"])))
self.ssl.to(get_torch_device()).eval()
self.dtm = DTM()
self.dtm.load_state_dict(torch.load(load_file_from_github_release(MODEL_TYPE, model_file_names["dtm"])))
self.dtm.to(get_torch_device()).eval()
def forward(self, img0, img1, t):
with torch.no_grad():
flow0, _ = self.raft(img0, img1)
flow1, _ = self.raft(img1, img0)
x = {
"images": torch.stack([img0, img1], dim=1),
"flows": torch.stack([flow0, flow1], dim=1),
}
out_ssl, _ = self.ssl(x, t=t, return_more=True)
out_dtm, _ = self.dtm(x, out_ssl, _, return_more=False)
return out_dtm[:, :3]
class EISAI_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (["eisai"], ),
"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: InterpolationStateList = None,
**kwargs
):
interpolation_model = EISAI(MODEL_FILE_NAMES)
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model):
return model(frame_0, frame_1, t=timestep)
scale = 1
args = [interpolation_model, scale]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
File diff suppressed because it is too large Load Diff
@@ -3,7 +3,7 @@ from comfy.model_management import get_torch_device, soft_empty_cache
import bisect
import numpy as np
import typing
from vfi_utils import InterpolationStateList, load_file_from_github_release, preprocess_frames, postprocess_frames
from vfi_utils import InterpolationStateListImport, load_file_from_github_release, preprocess_frames, postprocess_frames
import pathlib
import gc
@@ -41,7 +41,7 @@ def inference(model, img_batch_1, img_batch_2, inter_frames):
return [tensor.flip(0) for tensor in results]
class FILM_VFI:
class FILM_VFIImport:
@classmethod
def INPUT_TYPES(s):
return {
@@ -66,7 +66,7 @@ class FILM_VFI:
frames: torch.Tensor,
clear_cache_after_n_frames = 10,
multiplier: typing.SupportsInt = 2,
optional_interpolation_states: InterpolationStateList = None,
optional_interpolation_states: InterpolationStateListImport = None,
**kwargs
):
interpolation_states = optional_interpolation_states
@@ -1,115 +0,0 @@
import torch
from comfy.model_management import get_torch_device, soft_empty_cache
import numpy as np
import typing
from vfi_utils import InterpolationStateList, load_file_from_github_release, preprocess_frames, postprocess_frames, assert_batch_size
import pathlib
import warnings
from .flavr_arch import UNet_3D_3D, InputPadder
import gc
device = get_torch_device()
NBR_FRAME = 4
def build_flavr(model_path):
sd = torch.load(model_path)['state_dict']
sd = {k.partition("module.")[-1]:v for k,v in sd.items()}
#Ref: Class UNet_3D_3D
model = UNet_3D_3D("unet_18", n_inputs=NBR_FRAME, n_outputs=sd["outconv.1.weight"].shape[0] // 3, joinType="concat" , upmode="transpose")
model.load_state_dict(sd)
model.to(device).eval()
del sd
return model
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["FLAVR_2x.pth", "FLAVR_4x.pth", "FLAVR_8x.pth"]
class FLAVR_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 2, "max": 2}), #TODO: Implement recursively invoking interpolator for multi-frame interpolation
"duplicate_first_last_frames": ("BOOLEAN", {"default": False})
},
"optional": {
"optional_interpolation_states": ("INTERPOLATION_STATES", )
}
}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "vfi"
CATEGORY = "ComfyUI-Frame-Interpolation/VFI"
#Reference: https://github.com/danier97/ST-MFNet/blob/main/interpolate_yuv.py#L93
def vfi(
self,
ckpt_name: typing.AnyStr,
frames: torch.Tensor,
clear_cache_after_n_frames = 10,
multiplier: typing.SupportsInt = 2,
duplicate_first_last_frames: bool = False,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
if multiplier != 2:
warnings.warn("Currently, FLAVR only supports 2x interpolation. The process will continue but please set multiplier=2 afterward")
assert_batch_size(frames, batch_size=4, vfi_name="ST-MFNet")
interpolation_states = optional_interpolation_states
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
model = build_flavr(model_path)
frames = preprocess_frames(frames)
padder = InputPadder(frames.shape, 16)
frames = padder.pad(frames)
number_of_frames_processed_since_last_cleared_cuda_cache = 0
output_frames = []
for frame_itr in range(len(frames) - 3):
#Does skipping frame i+1 make sanse in this case?
if interpolation_states is not None and interpolation_states.is_frame_skipped(frame_itr) and interpolation_states.is_frame_skipped(frame_itr + 1):
continue
#Ensure that input frames are in fp32 - the same dtype as model
frame0, frame1, frame2, frame3 = (
frames[frame_itr:frame_itr+1].float(),
frames[frame_itr+1:frame_itr+2].float(),
frames[frame_itr+2:frame_itr+3].float(),
frames[frame_itr+3:frame_itr+4].float()
)
new_frame = model([frame0.to(device), frame1.to(device), frame2.to(device), frame3.to(device)])[0].detach().cpu()
number_of_frames_processed_since_last_cleared_cuda_cache += 2
if frame_itr == 0:
output_frames.append(frame0)
if duplicate_first_last_frames:
output_frames.append(frame0) # repeat the first frame
output_frames.append(frame1)
output_frames.append(new_frame)
output_frames.append(frame2)
if frame_itr == len(frames) - 4:
output_frames.append(frame3)
if duplicate_first_last_frames:
output_frames.append(frame3) # repeat the last frame
# 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()
dtype = torch.float32
output_frames = [frame.cpu().to(dtype=dtype) for frame in output_frames] #Ensure all frames are in cpu
out = torch.cat(output_frames, dim=0)
out = padder.unpad(out)
# clear cache for courtesy
print("Comfy-VFI: Final clearing cache...", end=' ')
soft_empty_cache()
print("Done cache clearing")
return (postprocess_frames(out), )
@@ -1,217 +0,0 @@
"""
https://github.com/tarun005/FLAVR/blob/main/model/FLAVR_arch.py
https://github.com/tarun005/FLAVR/blob/main/model/resnet_3D.py (only SEGating)
"""
import math
import numpy as np
import importlib
import torch
import torch.nn as nn
import torch.nn.functional as F
class SEGating(nn.Module):
def __init__(self , inplanes , reduction=16):
super().__init__()
self.pool = nn.AdaptiveAvgPool3d(1)
self.attn_layer = nn.Sequential(
nn.Conv3d(inplanes , inplanes , kernel_size=1 , stride=1 , bias=True),
nn.Sigmoid()
)
def forward(self , x):
out = self.pool(x)
y = self.attn_layer(out)
return x * y
def joinTensors(X1 , X2 , type="concat"):
if type == "concat":
return torch.cat([X1 , X2] , dim=1)
elif type == "add":
return X1 + X2
else:
return X1
class Conv_2d(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size, stride=1, padding=0, bias=False, batchnorm=False):
super().__init__()
self.conv = [nn.Conv2d(in_ch, out_ch, kernel_size=kernel_size, stride=stride, padding=padding, bias=bias)]
if batchnorm:
self.conv += [nn.BatchNorm2d(out_ch)]
self.conv = nn.Sequential(*self.conv)
def forward(self, x):
return self.conv(x)
class upConv3D(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size, stride, padding, upmode="transpose" , batchnorm=False):
super().__init__()
self.upmode = upmode
if self.upmode=="transpose":
self.upconv = nn.ModuleList(
[nn.ConvTranspose3d(in_ch, out_ch, kernel_size=kernel_size, stride=stride, padding=padding),
SEGating(out_ch)
]
)
else:
self.upconv = nn.ModuleList(
[nn.Upsample(mode='trilinear', scale_factor=(1,2,2), align_corners=False),
nn.Conv3d(in_ch, out_ch , kernel_size=1 , stride=1),
SEGating(out_ch)
]
)
if batchnorm:
self.upconv += [nn.BatchNorm3d(out_ch)]
self.upconv = nn.Sequential(*self.upconv)
def forward(self, x):
return self.upconv(x)
class Conv_3d(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size, stride=1, padding=0, bias=True, batchnorm=False):
super().__init__()
self.conv = [nn.Conv3d(in_ch, out_ch, kernel_size=kernel_size, stride=stride, padding=padding, bias=bias),
SEGating(out_ch)
]
if batchnorm:
self.conv += [nn.BatchNorm3d(out_ch)]
self.conv = nn.Sequential(*self.conv)
def forward(self, x):
return self.conv(x)
class upConv2D(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size, stride, padding, upmode="transpose" , batchnorm=False):
super().__init__()
self.upmode = upmode
if self.upmode=="transpose":
self.upconv = [nn.ConvTranspose2d(in_ch, out_ch, kernel_size=kernel_size, stride=stride, padding=padding)]
else:
self.upconv = [
nn.Upsample(mode='bilinear', scale_factor=2, align_corners=False),
nn.Conv2d(in_ch, out_ch , kernel_size=1 , stride=1)
]
if batchnorm:
self.upconv += [nn.BatchNorm2d(out_ch)]
self.upconv = nn.Sequential(*self.upconv)
def forward(self, x):
return self.upconv(x)
class UNet_3D_3D(nn.Module):
def __init__(self, block , n_inputs, n_outputs, batchnorm=False , joinType="concat" , upmode="transpose"):
super().__init__()
nf = [512 , 256 , 128 , 64]
out_channels = 3*n_outputs
self.joinType = joinType
self.n_outputs = n_outputs
growth = 2 if joinType == "concat" else 1
self.lrelu = nn.LeakyReLU(0.2, True)
unet_3D = importlib.import_module(".resnet_3D", "models.flavr")
if n_outputs > 1:
unet_3D.useBias = True
self.encoder = getattr(unet_3D , block)(pretrained=False , bn=batchnorm)
self.decoder = nn.Sequential(
Conv_3d(nf[0], nf[1] , kernel_size=3, padding=1, bias=True, batchnorm=batchnorm),
upConv3D(nf[1]*growth, nf[2], kernel_size=(3,4,4), stride=(1,2,2), padding=(1,1,1) , upmode=upmode, batchnorm=batchnorm),
upConv3D(nf[2]*growth, nf[3], kernel_size=(3,4,4), stride=(1,2,2), padding=(1,1,1) , upmode=upmode, batchnorm=batchnorm),
Conv_3d(nf[3]*growth, nf[3] , kernel_size=3, padding=1, bias=True, batchnorm=batchnorm),
upConv3D(nf[3]*growth , nf[3], kernel_size=(3,4,4), stride=(1,2,2), padding=(1,1,1) , upmode=upmode, batchnorm=batchnorm)
)
self.feature_fuse = Conv_2d(nf[3]*n_inputs , nf[3] , kernel_size=1 , stride=1, batchnorm=batchnorm)
self.outconv = nn.Sequential(
nn.ReflectionPad2d(3),
nn.Conv2d(nf[3], out_channels , kernel_size=7 , stride=1, padding=0)
)
def forward(self, images):
images = torch.stack(images , dim=2)
## Batch mean normalization works slightly better than global mean normalization, thanks to https://github.com/myungsub/CAIN
mean_ = images.mean(2, keepdim=True).mean(3, keepdim=True).mean(4,keepdim=True)
images = images-mean_
x_0 , x_1 , x_2 , x_3 , x_4 = self.encoder(images)
dx_3 = self.lrelu(self.decoder[0](x_4))
dx_3 = joinTensors(dx_3 , x_3 , type=self.joinType)
dx_2 = self.lrelu(self.decoder[1](dx_3))
dx_2 = joinTensors(dx_2 , x_2 , type=self.joinType)
dx_1 = self.lrelu(self.decoder[2](dx_2))
dx_1 = joinTensors(dx_1 , x_1 , type=self.joinType)
dx_0 = self.lrelu(self.decoder[3](dx_1))
dx_0 = joinTensors(dx_0 , x_0 , type=self.joinType)
dx_out = self.lrelu(self.decoder[4](dx_0))
dx_out = torch.cat(torch.unbind(dx_out , 2) , 1)
out = self.lrelu(self.feature_fuse(dx_out))
out = self.outconv(out)
out = torch.split(out, dim=1, split_size_or_sections=3)
mean_ = mean_.squeeze(2)
out = [o+mean_ for o in out]
return out
class InputPadder:
""" Pads images such that dimensions are divisible by divisor """
def __init__(self, dims, divisor=16):
self.ht, self.wd = dims[-2:]
pad_ht = (((self.ht // divisor) + 1) * divisor - self.ht) % divisor
pad_wd = (((self.wd // divisor) + 1) * divisor - self.wd) % divisor
self._pad = [pad_wd//2, pad_wd - pad_wd//2, pad_ht//2, pad_ht - pad_ht//2]
def pad(self, input_tensor):
return F.pad(input_tensor, self._pad, mode='replicate')
def unpad(self, input_tensor):
return self._unpad(input_tensor)
def _unpad(self, x):
ht, wd = x.shape[-2:]
c = [self._pad[2], ht-self._pad[3], self._pad[0], wd-self._pad[1]]
return x[..., c[0]:c[1], c[2]:c[3]]
@@ -1,288 +0,0 @@
# Modified from https://github.com/pytorch/vision/tree/master/torchvision/models/video
import torch
import torch.nn as nn
__all__ = ['unet_18', 'unet_34']
useBias = False
class identity(nn.Module):
def __init__(self , *args , **kwargs):
super().__init__()
def forward(self , x):
return x
class Conv3DSimple(nn.Conv3d):
def __init__(self,
in_planes,
out_planes,
midplanes=None,
stride=1,
padding=1):
super(Conv3DSimple, self).__init__(
in_channels=in_planes,
out_channels=out_planes,
kernel_size=(3, 3, 3),
stride=stride,
padding=padding,
bias=useBias)
@staticmethod
def get_downsample_stride(stride , temporal_stride):
if temporal_stride:
return (temporal_stride, stride, stride)
else:
return (stride , stride , stride)
class BasicStem(nn.Sequential):
"""The default conv-batchnorm-relu stem
"""
def __init__(self):
super().__init__(
nn.Conv3d(3, 64, kernel_size=(3, 7, 7), stride=(1, 2, 2),
padding=(1, 3, 3), bias=useBias),
batchnorm(64),
nn.ReLU(inplace=False))
class Conv2Plus1D(nn.Sequential):
def __init__(self,
in_planes,
out_planes,
midplanes,
stride=1,
padding=1):
if not isinstance(stride , int):
temporal_stride , stride , stride = stride
else:
temporal_stride = stride
super(Conv2Plus1D, self).__init__(
nn.Conv3d(in_planes, midplanes, kernel_size=(1, 3, 3),
stride=(1, stride, stride), padding=(0, padding, padding),
bias=False),
# batchnorm(midplanes),
nn.ReLU(inplace=True),
nn.Conv3d(midplanes, out_planes, kernel_size=(3, 1, 1),
stride=(temporal_stride, 1, 1), padding=(padding, 0, 0),
bias=False))
@staticmethod
def get_downsample_stride(stride , temporal_stride):
if temporal_stride:
return (temporal_stride, stride, stride)
else:
return (stride , stride , stride)
class R2Plus1dStem(nn.Sequential):
"""R(2+1)D stem is different than the default one as it uses separated 3D convolution
"""
def __init__(self):
super().__init__(
nn.Conv3d(3, 45, kernel_size=(1, 7, 7),
stride=(1, 2, 2), padding=(0, 3, 3),
bias=False),
batchnorm(45),
nn.ReLU(inplace=True),
nn.Conv3d(45, 64, kernel_size=(3, 1, 1),
stride=(1, 1, 1), padding=(1, 0, 0),
bias=False),
batchnorm(64),
nn.ReLU(inplace=True))
class SEGating(nn.Module):
def __init__(self , inplanes , reduction=16):
super().__init__()
self.pool = nn.AdaptiveAvgPool3d(1)
self.attn_layer = nn.Sequential(
nn.Conv3d(inplanes , inplanes , kernel_size=1 , stride=1 , bias=True),
nn.Sigmoid()
)
def forward(self , x):
out = self.pool(x)
y = self.attn_layer(out)
return x * y
class BasicBlock(nn.Module):
expansion = 1
def __init__(self, inplanes, planes, conv_builder, stride=1, downsample=None):
midplanes = (inplanes * planes * 3 * 3 * 3) // (inplanes * 3 * 3 + 3 * planes)
super(BasicBlock, self).__init__()
self.conv1 = nn.Sequential(
conv_builder(inplanes, planes, midplanes, stride),
batchnorm(planes),
nn.ReLU(inplace=True)
)
self.conv2 = nn.Sequential(
conv_builder(planes, planes, midplanes),
batchnorm(planes)
)
self.fg = SEGating(planes) ## Feature Gating
self.relu = nn.ReLU(inplace=True)
self.downsample = downsample
self.stride = stride
def forward(self, x):
residual = x
out = self.conv1(x)
out = self.conv2(out)
out = self.fg(out)
if self.downsample is not None:
residual = self.downsample(x)
out += residual
out = self.relu(out)
return out
class VideoResNet(nn.Module):
def __init__(self, block, conv_makers, layers,
stem, zero_init_residual=False):
"""Generic resnet video generator.
Args:
block (nn.Module): resnet building block
conv_makers (list(functions)): generator function for each layer
layers (List[int]): number of blocks per layer
stem (nn.Module, optional): Resnet stem, if None, defaults to conv-bn-relu. Defaults to None.
"""
super(VideoResNet, self).__init__()
self.inplanes = 64
self.stem = stem()
self.layer1 = self._make_layer(block, conv_makers[0], 64, layers[0], stride=1 )
self.layer2 = self._make_layer(block, conv_makers[1], 128, layers[1], stride=2 , temporal_stride=1)
self.layer3 = self._make_layer(block, conv_makers[2], 256, layers[2], stride=2 , temporal_stride=1)
self.layer4 = self._make_layer(block, conv_makers[3], 512, layers[3], stride=1, temporal_stride=1)
# init weights
self._initialize_weights()
if zero_init_residual:
for m in self.modules():
if isinstance(m, Bottleneck):
nn.init.constant_(m.bn3.weight, 0)
def forward(self, x):
x_0 = self.stem(x)
x_1 = self.layer1(x_0)
x_2 = self.layer2(x_1)
x_3 = self.layer3(x_2)
x_4 = self.layer4(x_3)
return x_0 , x_1 , x_2 , x_3 , x_4
def _make_layer(self, block, conv_builder, planes, blocks, stride=1, temporal_stride=None):
downsample = None
if stride != 1 or self.inplanes != planes * block.expansion:
ds_stride = conv_builder.get_downsample_stride(stride , temporal_stride)
downsample = nn.Sequential(
nn.Conv3d(self.inplanes, planes * block.expansion,
kernel_size=1, stride=ds_stride, bias=False),
batchnorm(planes * block.expansion)
)
stride = ds_stride
layers = []
layers.append(block(self.inplanes, planes, conv_builder, stride, downsample ))
self.inplanes = planes * block.expansion
for i in range(1, blocks):
layers.append(block(self.inplanes, planes, conv_builder ))
return nn.Sequential(*layers)
def _initialize_weights(self):
for m in self.modules():
if isinstance(m, nn.Conv3d):
nn.init.kaiming_normal_(m.weight, mode='fan_out',
nonlinearity='relu')
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.BatchNorm3d):
nn.init.constant_(m.weight, 1)
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.Linear):
nn.init.normal_(m.weight, 0, 0.01)
nn.init.constant_(m.bias, 0)
def _video_resnet(arch, pretrained=False, progress=True, **kwargs):
model = VideoResNet(**kwargs)
## TODO: Other 3D resnet models, like S3D, r(2+1)D.
if pretrained:
state_dict = load_state_dict_from_url(model_urls[arch],
progress=progress)
model.load_state_dict(state_dict)
return model
def unet_18(pretrained=False, bn=False, progress=True, **kwargs):
"""
Construct 18 layer Unet3D model as in
https://arxiv.org/abs/1711.11248
Args:
pretrained (bool): If True, returns a model pre-trained on Kinetics-400
progress (bool): If True, displays a progress bar of the download to stderr
Returns:
nn.Module: R3D-18 encoder
"""
global batchnorm
if bn:
batchnorm = nn.BatchNorm3d
else:
batchnorm = identity
return _video_resnet('r3d_18',
pretrained, progress,
block=BasicBlock,
conv_makers=[Conv3DSimple] * 4,
layers=[2, 2, 2, 2],
stem=BasicStem, **kwargs)
def unet_34(pretrained=False, bn=False, progress=True, **kwargs):
"""
Construct 34 layer Unet3D model as in
https://arxiv.org/abs/1711.11248
Args:
pretrained (bool): If True, returns a model pre-trained on Kinetics-400
progress (bool): If True, displays a progress bar of the download to stderr
Returns:
nn.Module: R3D-18 encoder
"""
global batchnorm
# bn = False
if bn:
batchnorm = nn.BatchNorm3d
else:
batchnorm = identity
return _video_resnet('r3d_34',
pretrained, progress,
block=BasicBlock,
conv_makers=[Conv3DSimple] * 4,
layers=[3, 4, 6, 3],
stem=BasicStem, **kwargs)
@@ -1,24 +0,0 @@
import itertools
import numpy as np
import vapoursynth as vs
from .GMFSS_Fortuna_arch import Model_inference
import torch
import traceback
class GMFSS_Fortuna:
def __init__(self):
self.cache = False
self.amount_input_img = 2
torch.set_grad_enabled(False)
torch.backends.cudnn.enabled = True
torch.backends.cudnn.benchmark = True
self.model = Model_inference()
self.model.eval()
def execute(self, I0, I1, timestep):
with torch.inference_mode():
middle = self.model(I0, I1, timestep).cpu()
return middle
@@ -1,23 +0,0 @@
import itertools
import numpy as np
import vapoursynth as vs
from .GMFSS_Fortuna_union_arch import Model_inference
import torch
class GMFSS_Fortuna_union:
def __init__(self):
self.cache = False
self.amount_input_img = 2
torch.set_grad_enabled(False)
torch.backends.cudnn.enabled = True
torch.backends.cudnn.benchmark = True
self.model = Model_inference()
self.model.eval()
def execute(self, I0, I1, timestep):
with torch.inference_mode():
middle = self.model(I0, I1, timestep).cpu()
return middle
@@ -1,143 +0,0 @@
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
import torch
import torch.nn as nn
import torch.nn.functional as F
from comfy.model_management import get_torch_device
GLOBAL_MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPTS_PATH_CONFIG = {
"GMFSS_fortuna_union": {
"ifnet": ("rife", "rife46.pth"),
"flownet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_flownet.pkl"),
"metricnet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_union_metric.pkl"),
"feat_ext": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_union_feat.pkl"),
"fusionnet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_union_fusionnet.pkl")
},
"GMFSS_fortuna": {
"flownet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_flownet.pkl"),
"metricnet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_metric.pkl"),
"feat_ext": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_feat.pkl"),
"fusionnet": (GLOBAL_MODEL_TYPE, "GMFSS_fortuna_fusionnet.pkl")
}
}
class CommonModelInference(nn.Module):
def __init__(self, model_type):
super(CommonModelInference, self).__init__()
from .GMFSS_Fortuna_arch import Model as GMFSS
from .GMFSS_Fortuna_union_arch import Model as GMFSS_Union
self.model = GMFSS_Union() if "union" in model_type else GMFSS()
self.model.eval()
self.model.device()
_model_path_config = CKPTS_PATH_CONFIG[model_type]
self.model.load_model({
key: load_file_from_github_release(*_model_path_config[key])
for key in _model_path_config
})
def forward(self, I0, I1, timestep, scale=1.0):
n, c, h, w = I0.shape
tmp = max(64, int(64 / scale))
ph = ((h - 1) // tmp + 1) * tmp
pw = ((w - 1) // tmp + 1) * tmp
padding = (0, pw - w, 0, ph - h)
I0 = F.pad(I0, padding)
I1 = F.pad(I1, padding)
(
flow01,
flow10,
metric0,
metric1,
feat11,
feat12,
feat13,
feat21,
feat22,
feat23,
) = self.model.reuse(I0, I1, scale)
output = self.model.inference(
I0,
I1,
flow01,
flow10,
metric0,
metric1,
feat11,
feat12,
feat13,
feat21,
feat22,
feat23,
timestep,
)
return output[:, :, :h, :w]
class GMFSS_Fortuna_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (list(CKPTS_PATH_CONFIG.keys()), ),
"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: InterpolationStateList = None,
**kwargs
):
"""
Perform video frame interpolation using a given checkpoint model.
Args:
ckpt_name (str): The name of the checkpoint model to use.
frames (torch.Tensor): A tensor containing input video frames.
clear_cache_after_n_frames (int, optional): The number of frames to process before clearing CUDA cache
to prevent memory overflow. Defaults to 10. Lower numbers are safer but mean more processing time.
How high you should set it depends on how many input frames there are, input resolution (after upscaling),
how many times you want to multiply them, and how long you're willing to wait for the process to complete.
multiplier (int, optional): The multiplier for each input frame. 60 input frames * 2 = 120 output frames. Defaults to 2.
Returns:
tuple: A tuple containing the output interpolated frames.
Note:
This method interpolates frames in a video sequence using a specified checkpoint model.
It processes each frame sequentially, generating interpolated frames between them.
To prevent memory overflow, it clears the CUDA cache after processing a specified number of frames.
"""
interpolation_model = CommonModelInference(model_type=ckpt_name)
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model, scale):
return model(frame_0, frame_1, timestep, scale)
scale = 1
args = [interpolation_model, scale]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
@@ -1,293 +0,0 @@
# https://github.com/ltkong218/IFRNet/blob/main/models/IFRNet_L.py
# https://github.com/ltkong218/IFRNet/blob/main/utils.py
import torch
import torch.nn as nn
import torch.nn.functional as F
from comfy.model_management import get_torch_device
def warp(img, flow):
B, _, H, W = flow.shape
xx = torch.linspace(-1.0, 1.0, W).view(1, 1, 1, W).expand(B, -1, H, -1)
yy = torch.linspace(-1.0, 1.0, H).view(1, 1, H, 1).expand(B, -1, -1, W)
grid = torch.cat([xx, yy], 1).to(img)
flow_ = torch.cat(
[
flow[:, 0:1, :, :] / ((W - 1.0) / 2.0),
flow[:, 1:2, :, :] / ((H - 1.0) / 2.0),
],
1,
)
grid_ = (grid + flow_).permute(0, 2, 3, 1)
output = F.grid_sample(
input=img,
grid=grid_,
mode="bilinear",
padding_mode="border",
align_corners=True,
)
return output
def get_robust_weight(flow_pred, flow_gt, beta):
epe = ((flow_pred.detach() - flow_gt) ** 2).sum(dim=1, keepdim=True) ** 0.5
robust_weight = torch.exp(-beta * epe)
return robust_weight
def resize(x, scale_factor):
return F.interpolate(
x, scale_factor=scale_factor, mode="bilinear", align_corners=False
)
def convrelu(
in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
groups=1,
bias=True,
):
return nn.Sequential(
nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride,
padding,
dilation,
groups,
bias=bias,
),
nn.PReLU(out_channels),
)
class ResBlock(nn.Module):
def __init__(self, in_channels, side_channels, bias=True):
super(ResBlock, self).__init__()
self.side_channels = side_channels
self.conv1 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv2 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv3 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv4 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv5 = nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
)
self.prelu = nn.PReLU(in_channels)
def forward(self, x):
out = self.conv1(x)
out[:, -self.side_channels :, :, :] = self.conv2(
out[:, -self.side_channels :, :, :]
)
out = self.conv3(out)
out[:, -self.side_channels :, :, :] = self.conv4(
out[:, -self.side_channels :, :, :]
)
out = self.prelu(x + self.conv5(out))
return out
class Encoder(nn.Module):
def __init__(self):
super(Encoder, self).__init__()
self.pyramid1 = nn.Sequential(
convrelu(3, 64, 7, 2, 3), convrelu(64, 64, 3, 1, 1)
)
self.pyramid2 = nn.Sequential(
convrelu(64, 96, 3, 2, 1), convrelu(96, 96, 3, 1, 1)
)
self.pyramid3 = nn.Sequential(
convrelu(96, 144, 3, 2, 1), convrelu(144, 144, 3, 1, 1)
)
self.pyramid4 = nn.Sequential(
convrelu(144, 192, 3, 2, 1), convrelu(192, 192, 3, 1, 1)
)
def forward(self, img):
f1 = self.pyramid1(img)
f2 = self.pyramid2(f1)
f3 = self.pyramid3(f2)
f4 = self.pyramid4(f3)
return f1, f2, f3, f4
class Decoder4(nn.Module):
def __init__(self):
super(Decoder4, self).__init__()
self.convblock = nn.Sequential(
convrelu(384 + 1, 384),
ResBlock(384, 64),
nn.ConvTranspose2d(384, 148, 4, 2, 1, bias=True),
)
def forward(self, f0, f1, embt):
b, c, h, w = f0.shape
embt = embt.repeat(1, 1, h, w)
f_in = torch.cat([f0, f1, embt], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder3(nn.Module):
def __init__(self):
super(Decoder3, self).__init__()
self.convblock = nn.Sequential(
convrelu(436, 432),
ResBlock(432, 64),
nn.ConvTranspose2d(432, 100, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder2(nn.Module):
def __init__(self):
super(Decoder2, self).__init__()
self.convblock = nn.Sequential(
convrelu(292, 288),
ResBlock(288, 64),
nn.ConvTranspose2d(288, 68, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder1(nn.Module):
def __init__(self):
super(Decoder1, self).__init__()
self.convblock = nn.Sequential(
convrelu(196, 192),
ResBlock(192, 64),
nn.ConvTranspose2d(192, 8, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class IRFNet_L(nn.Module):
def __init__(self):
super(IRFNet_L, self).__init__()
self.encoder = Encoder()
self.decoder4 = Decoder4()
self.decoder3 = Decoder3()
self.decoder2 = Decoder2()
self.decoder1 = Decoder1()
def forward(self, img0, img1, scale_factor=1.0, timestep=0.5):
# emb1 = torch.tensor(1/2).view(1, 1, 1, 1).float()
# emb2 = torch.tensor(2/2).view(1, 1, 1, 1).float()
# embt = torch.cat([emb1, emb2], 0)
n, c, h, w = img0.shape
ph = ((h - 1) // 64 + 1) * 64
pw = ((w - 1) // 64 + 1) * 64
padding = (0, pw - w, 0, ph - h)
img0 = F.pad(img0, padding)
img1 = F.pad(img1, padding)
#Support multiple batches
embt = torch.tensor([timestep] * n).view(n, 1, 1, 1).float().to(get_torch_device())
if "HalfTensor" in str(img0.type()):
embt = embt.half()
mean_ = (
torch.cat([img0, img1], 2)
.mean(1, keepdim=True)
.mean(2, keepdim=True)
.mean(3, keepdim=True)
)
img0 = img0 - mean_
img1 = img1 - mean_
img0_ = resize(img0, scale_factor=scale_factor)
img1_ = resize(img1, scale_factor=scale_factor)
f0_1, f0_2, f0_3, f0_4 = self.encoder(img0_)
f1_1, f1_2, f1_3, f1_4 = self.encoder(img1_)
out4 = self.decoder4(f0_4, f1_4, embt)
up_flow0_4 = out4[:, 0:2]
up_flow1_4 = out4[:, 2:4]
ft_3_ = out4[:, 4:]
out3 = self.decoder3(ft_3_, f0_3, f1_3, up_flow0_4, up_flow1_4)
up_flow0_3 = out3[:, 0:2] + 2.0 * resize(up_flow0_4, scale_factor=2.0)
up_flow1_3 = out3[:, 2:4] + 2.0 * resize(up_flow1_4, scale_factor=2.0)
ft_2_ = out3[:, 4:]
out2 = self.decoder2(ft_2_, f0_2, f1_2, up_flow0_3, up_flow1_3)
up_flow0_2 = out2[:, 0:2] + 2.0 * resize(up_flow0_3, scale_factor=2.0)
up_flow1_2 = out2[:, 2:4] + 2.0 * resize(up_flow1_3, scale_factor=2.0)
ft_1_ = out2[:, 4:]
out1 = self.decoder1(ft_1_, f0_1, f1_1, up_flow0_2, up_flow1_2)
up_flow0_1 = out1[:, 0:2] + 2.0 * resize(up_flow0_2, scale_factor=2.0)
up_flow1_1 = out1[:, 2:4] + 2.0 * resize(up_flow1_2, scale_factor=2.0)
up_mask_1 = torch.sigmoid(out1[:, 4:5])
up_res_1 = out1[:, 5:]
up_flow0_1 = resize(up_flow0_1, scale_factor=(1.0 / scale_factor)) * (
1.0 / scale_factor
)
up_flow1_1 = resize(up_flow1_1, scale_factor=(1.0 / scale_factor)) * (
1.0 / scale_factor
)
up_mask_1 = resize(up_mask_1, scale_factor=(1.0 / scale_factor))
up_res_1 = resize(up_res_1, scale_factor=(1.0 / scale_factor))
img0_warp = warp(img0, up_flow0_1)
img1_warp = warp(img1, up_flow1_1)
imgt_merge = up_mask_1 * img0_warp + (1 - up_mask_1) * img1_warp + mean_
imgt_pred = imgt_merge + up_res_1
imgt_pred = torch.clamp(imgt_pred, 0, 1)
return imgt_pred[:, :, :h, :w]
@@ -1,293 +0,0 @@
# https://github.com/ltkong218/IFRNet/blob/main/models/IFRNet_S.py
# https://github.com/ltkong218/IFRNet/blob/main/utils.py
import torch
import torch.nn as nn
import torch.nn.functional as F
from comfy.model_management import get_torch_device
def warp(img, flow):
B, _, H, W = flow.shape
xx = torch.linspace(-1.0, 1.0, W).view(1, 1, 1, W).expand(B, -1, H, -1)
yy = torch.linspace(-1.0, 1.0, H).view(1, 1, H, 1).expand(B, -1, -1, W)
grid = torch.cat([xx, yy], 1).to(img)
flow_ = torch.cat(
[
flow[:, 0:1, :, :] / ((W - 1.0) / 2.0),
flow[:, 1:2, :, :] / ((H - 1.0) / 2.0),
],
1,
)
grid_ = (grid + flow_).permute(0, 2, 3, 1)
output = F.grid_sample(
input=img,
grid=grid_,
mode="bilinear",
padding_mode="border",
align_corners=True,
)
return output
def get_robust_weight(flow_pred, flow_gt, beta):
epe = ((flow_pred.detach() - flow_gt) ** 2).sum(dim=1, keepdim=True) ** 0.5
robust_weight = torch.exp(-beta * epe)
return robust_weight
def resize(x, scale_factor):
return F.interpolate(
x, scale_factor=scale_factor, mode="bilinear", align_corners=False
)
def convrelu(
in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
groups=1,
bias=True,
):
return nn.Sequential(
nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride,
padding,
dilation,
groups,
bias=bias,
),
nn.PReLU(out_channels),
)
class ResBlock(nn.Module):
def __init__(self, in_channels, side_channels, bias=True):
super(ResBlock, self).__init__()
self.side_channels = side_channels
self.conv1 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv2 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv3 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv4 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv5 = nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
)
self.prelu = nn.PReLU(in_channels)
def forward(self, x):
out = self.conv1(x)
out[:, -self.side_channels :, :, :] = self.conv2(
out[:, -self.side_channels :, :, :]
)
out = self.conv3(out)
out[:, -self.side_channels :, :, :] = self.conv4(
out[:, -self.side_channels :, :, :]
)
out = self.prelu(x + self.conv5(out))
return out
class Encoder(nn.Module):
def __init__(self):
super(Encoder, self).__init__()
self.pyramid1 = nn.Sequential(
convrelu(3, 24, 3, 2, 1), convrelu(24, 24, 3, 1, 1)
)
self.pyramid2 = nn.Sequential(
convrelu(24, 36, 3, 2, 1), convrelu(36, 36, 3, 1, 1)
)
self.pyramid3 = nn.Sequential(
convrelu(36, 54, 3, 2, 1), convrelu(54, 54, 3, 1, 1)
)
self.pyramid4 = nn.Sequential(
convrelu(54, 72, 3, 2, 1), convrelu(72, 72, 3, 1, 1)
)
def forward(self, img):
f1 = self.pyramid1(img)
f2 = self.pyramid2(f1)
f3 = self.pyramid3(f2)
f4 = self.pyramid4(f3)
return f1, f2, f3, f4
class Decoder4(nn.Module):
def __init__(self):
super(Decoder4, self).__init__()
self.convblock = nn.Sequential(
convrelu(144 + 1, 144),
ResBlock(144, 24),
nn.ConvTranspose2d(144, 58, 4, 2, 1, bias=True),
)
def forward(self, f0, f1, embt):
b, c, h, w = f0.shape
embt = embt.repeat(1, 1, h, w)
f_in = torch.cat([f0, f1, embt], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder3(nn.Module):
def __init__(self):
super(Decoder3, self).__init__()
self.convblock = nn.Sequential(
convrelu(166, 162),
ResBlock(162, 24),
nn.ConvTranspose2d(162, 40, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder2(nn.Module):
def __init__(self):
super(Decoder2, self).__init__()
self.convblock = nn.Sequential(
convrelu(112, 108),
ResBlock(108, 24),
nn.ConvTranspose2d(108, 28, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class Decoder1(nn.Module):
def __init__(self):
super(Decoder1, self).__init__()
self.convblock = nn.Sequential(
convrelu(76, 72),
ResBlock(72, 24),
nn.ConvTranspose2d(72, 8, 4, 2, 1, bias=True),
)
def forward(self, ft_, f0, f1, up_flow0, up_flow1):
f0_warp = warp(f0, up_flow0)
f1_warp = warp(f1, up_flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, up_flow0, up_flow1], 1)
f_out = self.convblock(f_in)
return f_out
class IRFNet_S(nn.Module):
def __init__(self):
super(IRFNet_S, self).__init__()
self.encoder = Encoder()
self.decoder4 = Decoder4()
self.decoder3 = Decoder3()
self.decoder2 = Decoder2()
self.decoder1 = Decoder1()
def forward(self, img0, img1, scale_factor=1.0, timestep=0.5):
# emb1 = torch.tensor(1/2).view(1, 1, 1, 1).float()
# emb2 = torch.tensor(2/2).view(1, 1, 1, 1).float()
# embt = torch.cat([emb1, emb2], 0)
n, c, h, w = img0.shape
ph = ((h - 1) // 64 + 1) * 64
pw = ((w - 1) // 64 + 1) * 64
padding = (0, pw - w, 0, ph - h)
img0 = F.pad(img0, padding)
img1 = F.pad(img1, padding)
#Support multiple batches
embt = torch.tensor([timestep] * n).view(n, 1, 1, 1).float().to(get_torch_device())
if "HalfTensor" in str(img0.type()):
embt = embt.half()
mean_ = (
torch.cat([img0, img1], 2)
.mean(1, keepdim=True)
.mean(2, keepdim=True)
.mean(3, keepdim=True)
)
img0 = img0 - mean_
img1 = img1 - mean_
img0_ = resize(img0, scale_factor=scale_factor)
img1_ = resize(img1, scale_factor=scale_factor)
f0_1, f0_2, f0_3, f0_4 = self.encoder(img0_)
f1_1, f1_2, f1_3, f1_4 = self.encoder(img1_)
out4 = self.decoder4(f0_4, f1_4, embt)
up_flow0_4 = out4[:, 0:2]
up_flow1_4 = out4[:, 2:4]
ft_3_ = out4[:, 4:]
out3 = self.decoder3(ft_3_, f0_3, f1_3, up_flow0_4, up_flow1_4)
up_flow0_3 = out3[:, 0:2] + 2.0 * resize(up_flow0_4, scale_factor=2.0)
up_flow1_3 = out3[:, 2:4] + 2.0 * resize(up_flow1_4, scale_factor=2.0)
ft_2_ = out3[:, 4:]
out2 = self.decoder2(ft_2_, f0_2, f1_2, up_flow0_3, up_flow1_3)
up_flow0_2 = out2[:, 0:2] + 2.0 * resize(up_flow0_3, scale_factor=2.0)
up_flow1_2 = out2[:, 2:4] + 2.0 * resize(up_flow1_3, scale_factor=2.0)
ft_1_ = out2[:, 4:]
out1 = self.decoder1(ft_1_, f0_1, f1_1, up_flow0_2, up_flow1_2)
up_flow0_1 = out1[:, 0:2] + 2.0 * resize(up_flow0_2, scale_factor=2.0)
up_flow1_1 = out1[:, 2:4] + 2.0 * resize(up_flow1_2, scale_factor=2.0)
up_mask_1 = torch.sigmoid(out1[:, 4:5])
up_res_1 = out1[:, 5:]
up_flow0_1 = resize(up_flow0_1, scale_factor=(1.0 / scale_factor)) * (
1.0 / scale_factor
)
up_flow1_1 = resize(up_flow1_1, scale_factor=(1.0 / scale_factor)) * (
1.0 / scale_factor
)
up_mask_1 = resize(up_mask_1, scale_factor=(1.0 / scale_factor))
up_res_1 = resize(up_res_1, scale_factor=(1.0 / scale_factor))
img0_warp = warp(img0, up_flow0_1)
img1_warp = warp(img1, up_flow1_1)
imgt_merge = up_mask_1 * img0_warp + (1 - up_mask_1) * img1_warp + mean_
imgt_pred = imgt_merge + up_res_1
imgt_pred = torch.clamp(imgt_pred, 0, 1)
return imgt_pred[:, :, :h, :w]
@@ -1,57 +0,0 @@
import torch
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames
import typing
from comfy.model_management import get_torch_device
from vfi_utils import generic_frame_loop, InterpolationStateList
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["IFRNet_S_Vimeo90K.pth", "IFRNet_L_Vimeo90K.pth"]
class IFRNet_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 2, "max": 1000}),
"scale_factor": ([0.25, 0.5, 1.0, 2.0, 4.0], {"default": 1.0}),
},
"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: typing.SupportsInt = 1,
multiplier: typing.SupportsInt = 2,
scale_factor: typing.SupportsFloat = 1.0,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
from .IFRNet_S_arch import IRFNet_S
from .IFRNet_L_arch import IRFNet_L
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
interpolation_model = IRFNet_S() if 'S' in ckpt_name else IRFNet_L()
interpolation_model.load_state_dict(torch.load(model_path))
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model, scale_factor):
return model(frame_0, frame_1, timestep, scale_factor)
args = [interpolation_model, scale_factor]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
@@ -1,766 +0,0 @@
"""
https://github.com/98mxr/IFUNet/blob/main/model/IFUNet.py
https://github.com/98mxr/IFUNet/blob/main/model/cbam.py
https://github.com/98mxr/IFUNet/blob/main/model/warplayer.py
https://github.com/98mxr/IFUNet/blob/5be535c8cff66d6fa1967252685719df4c0620e4/model/RIFE.py
https://github.com/98mxr/IFUNet/blob/main/model/rrdb.py
https://github.com/98mxr/IFUNet/blob/main/model/ResynNet.py
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from comfy.model_management import get_torch_device
backwarp_tenGrid = {}
device = get_torch_device()
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
),
nn.PReLU(out_planes),
)
def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=False,
),
nn.BatchNorm2d(out_planes),
nn.PReLU(out_planes),
)
class DegCNN(nn.Module):
def __init__(self):
super(DegCNN, self).__init__()
self.conv0 = conv(3, 32, 3, 2, 1)
self.conv1 = conv(32, 32, 3, 2, 1)
self.conv2 = conv(32, 32, 3, 2, 1)
self.conv3 = conv(32, 32, 3, 2, 1)
self.deconv = nn.Sequential(
nn.Dropout2d(0.95),
nn.ConvTranspose2d(4 * 32, 32, 4, 2, 1),
nn.PReLU(32),
nn.Conv2d(32, 3, 3, 1, 1),
nn.Sigmoid(),
)
def forward(self, x):
f0 = self.conv0(x)
f1 = self.conv1(f0)
f2 = self.conv2(f1)
f3 = self.conv3(f2)
f1 = F.interpolate(f1, scale_factor=2.0, mode="bilinear", align_corners=False)
f2 = F.interpolate(f2, scale_factor=4.0, mode="bilinear", align_corners=False)
f3 = F.interpolate(f3, scale_factor=8.0, mode="bilinear", align_corners=False)
return self.deconv(torch.cat((f0, f1, f2, f3), 1))
class FlowBlock(nn.Module):
def __init__(self, in_planes, c=64):
super(FlowBlock, self).__init__()
self.conv0 = nn.Sequential(
conv_bn(in_planes, c // 2, 3, 2, 1),
conv_bn(c // 2, c, 3, 2, 1),
conv_bn(c, 2 * c, 3, 2, 1),
)
self.convblock = nn.Sequential(
conv_bn(2 * c, 2 * c),
conv_bn(2 * c, 2 * c),
conv_bn(2 * c, 2 * c),
conv_bn(2 * c, 2 * c),
conv_bn(2 * c, 2 * c),
conv_bn(2 * c, 2 * c),
)
self.lastconv = nn.ConvTranspose2d(2 * c, 4, 4, 2, 1)
def forward(self, x, flow, scale=1):
x = F.interpolate(
x, scale_factor=1.0 / scale, mode="bilinear", align_corners=False
)
if flow is not None:
flow = (
F.interpolate(
flow, scale_factor=1.0 / scale, mode="bilinear", align_corners=False
)
* 1.0
/ scale
)
x = torch.cat((x, flow), 1)
feat = self.conv0(x)
feat = self.convblock(feat) + feat
tmp = self.lastconv(feat)
tmp = F.interpolate(
tmp, scale_factor=scale * 4, mode="bilinear", align_corners=False
)
flow = tmp[:, :2] * scale * 4
mask = tmp[:, 2:3]
return flow, mask
class ResynNet(nn.Module):
def __init__(self):
super(ResynNet, self).__init__()
self.block0 = FlowBlock(6, c=128)
self.block1 = FlowBlock(12, c=128)
self.block2 = FlowBlock(12, c=128)
self.degrad = DegCNN()
# Contextual Refinement context + decode
self.context0 = nn.Sequential(
conv(3, 16, 3, 2, 1),
conv(16, 32, 3, 2, 1),
)
self.context1 = nn.Sequential(
conv(3, 16, 3, 2, 1),
conv(16, 32, 3, 2, 1),
)
self.decode = nn.Sequential(
nn.ConvTranspose2d(64, 32, 4, 2, 1),
nn.ConvTranspose2d(32, 3, 4, 2, 1),
nn.Tanh(),
)
def calflow(self, img0, lowres, scale):
flow = None
stu = [self.block0, self.block1, self.block2]
for i in range(3):
if flow is not None:
flow_d, mask_d = stu[i](
torch.cat((img0, lowres, warped_img0, mask), 1),
flow,
scale=scale[i],
)
flow = flow + flow_d
mask = mask + mask_d
else:
flow, mask = stu[i](torch.cat((img0, lowres), 1), None, scale=scale[i])
warped_img0 = warp(img0, flow)
flow_down = (
F.interpolate(flow, scale_factor=0.25, mode="bilinear", align_corners=False)
* 0.25
)
c0 = warp(self.context0(img0), flow_down)
c1 = self.context1(warped_img0)
warped_img0 = warped_img0 + self.decode(torch.cat((c0, c1), 1))
return flow, mask, torch.clamp(warped_img0, 0, 1)
def forward(
self, x, deg=None, gt=None, scale=[4, 2, 1], training=False, blend=True
):
if training:
deg = self.degrad(gt)
loss_cons = (gt - deg).abs().mean()
else:
loss_cons = torch.tensor([0])
img_list = []
N = x.shape[1] // 3
for i in range(N):
img_list.append(x[:, i * 3 : i * 3 + 3])
warped_list = []
merged = []
mask_list = []
flow_list = []
for i in range(N):
f, m, img = self.calflow(img_list[i], deg.detach(), scale)
mask_list.append(m)
warped_list.append(img)
flow_list.append(f)
if blend:
N += 1
mask_list.append(m * 0)
warped_list.append(deg)
mask = F.softmax(torch.clamp(torch.cat(mask_list, 1), -4, 4), dim=1)
merged = 0
for i in range(N):
merged += warped_list[i] * mask[:, i : i + 1]
return merged, loss_cons
def make_layer(basic_block, num_basic_block, **kwarg):
"""Make layers by stacking the same blocks.
Args:
basic_block (nn.module): nn.module class for basic block.
num_basic_block (int): number of blocks.
Returns:
nn.Sequential: Stacked blocks in nn.Sequential.
"""
layers = []
for _ in range(num_basic_block):
layers.append(basic_block(**kwarg))
return nn.Sequential(*layers)
class ResidualDenseBlock(nn.Module):
"""Residual Dense Block.
Used in RRDB block in ESRGAN.
Args:
num_feat (int): Channel number of intermediate features.
num_grow_ch (int): Channels for each growth.
"""
def __init__(self, num_feat=64, num_grow_ch=32):
super(ResidualDenseBlock, self).__init__()
self.conv1 = nn.Conv2d(num_feat, num_grow_ch, 3, 1, 1)
self.conv2 = nn.Conv2d(num_feat + num_grow_ch, num_grow_ch, 3, 1, 1)
self.conv3 = nn.Conv2d(num_feat + 2 * num_grow_ch, num_grow_ch, 3, 1, 1)
self.conv4 = nn.Conv2d(num_feat + 3 * num_grow_ch, num_grow_ch, 3, 1, 1)
self.conv5 = nn.Conv2d(num_feat + 4 * num_grow_ch, num_feat, 3, 1, 1)
self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True)
# initialization
# default_init_weights([self.conv1, self.conv2, self.conv3, self.conv4, self.conv5], 0.1)
# 只能先取消,default_init_weights来自basicsr.arch_util
def forward(self, x):
x1 = self.lrelu(self.conv1(x))
x2 = self.lrelu(self.conv2(torch.cat((x, x1), 1)))
x3 = self.lrelu(self.conv3(torch.cat((x, x1, x2), 1)))
x4 = self.lrelu(self.conv4(torch.cat((x, x1, x2, x3), 1)))
x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1))
# Emperically, we use 0.2 to scale the residual for better performance
# 原作者这么说我就这么听着吧
return x5 * 0.2 + x
class RRDB(nn.Module):
"""Residual in Residual Dense Block.
Used in RRDB-Net in ESRGAN.
Args:
num_feat (int): Channel number of intermediate features.
num_grow_ch (int): Channels for each growth.
"""
def __init__(self, num_feat, num_grow_ch=32):
super(RRDB, self).__init__()
self.rdb1 = ResidualDenseBlock(num_feat, num_grow_ch)
self.rdb2 = ResidualDenseBlock(num_feat, num_grow_ch)
self.rdb3 = ResidualDenseBlock(num_feat, num_grow_ch)
def forward(self, x):
out = self.rdb1(x)
out = self.rdb2(out)
out = self.rdb3(out)
# Emperically, we use 0.2 to scale the residual for better performance
# 原作者这么说我就这么听着吧
return out * 0.2 + x
class RRDBNet(nn.Module):
"""Networks consisting of Residual in Residual Dense Block, which is used
in ESRGAN.
ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks.
We extend ESRGAN for scale x2 and scale x1.
Note: This is one option for scale 1, scale 2 in RRDBNet.
We first employ the pixel-unshuffle (an inverse operation of pixelshuffle to reduce the spatial size
and enlarge the channel size before feeding inputs into the main ESRGAN architecture.
Args:
num_in_ch (int): Channel number of inputs.
num_out_ch (int): Channel number of outputs.
num_feat (int): Channel number of intermediate features.
Default: 64
num_block (int): Block number in the trunk network. Defaults: 23
num_grow_ch (int): Channels for each growth. Default: 32.
"""
def __init__(
self, num_in_ch=16, num_out_ch=1, num_feat=64, num_block=6, num_grow_ch=32
):
super(RRDBNet, self).__init__()
self.conv_first = nn.Conv2d(num_in_ch, num_feat, 3, 1, 1)
self.body = make_layer(
RRDB, num_block, num_feat=num_feat, num_grow_ch=num_grow_ch
)
self.conv_body = nn.Conv2d(num_feat, num_feat, 3, 1, 1)
# upsample
self.conv_up1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1)
self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1)
self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1)
self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1)
self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True)
def forward(self, img0, img1, warped_img0, warped_img1, flow):
x = torch.cat((img0, img1, warped_img0, warped_img1), 1)
x = F.interpolate(x, scale_factor=0.25, mode="bilinear", align_corners=False)
flow = (
F.interpolate(flow, scale_factor=0.25, mode="bilinear", align_corners=False)
* 0.25
)
feat = torch.cat((x, flow), 1)
feat = self.conv_first(feat)
body_feat = self.conv_body(self.body(feat))
feat = feat + body_feat
# upsample,充分利用四倍放大
feat = self.lrelu(
self.conv_up1(F.interpolate(feat, scale_factor=2.0, mode="nearest"))
)
feat = self.lrelu(
self.conv_up2(F.interpolate(feat, scale_factor=2.0, mode="nearest"))
)
out = self.conv_last(self.lrelu(self.conv_hr(feat)))
out = torch.sigmoid(out)
return out
def warp(tenInput, tenFlow):
k = (str(tenFlow.device), str(tenFlow.size()))
if k not in backwarp_tenGrid:
tenHorizontal = (
torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device)
.view(1, 1, 1, tenFlow.shape[3])
.expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)
)
tenVertical = (
torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device)
.view(1, 1, tenFlow.shape[2], 1)
.expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])
)
backwarp_tenGrid[k] = torch.cat([tenHorizontal, tenVertical], 1).to(device)
tenFlow = torch.cat(
[
tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),
tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0),
],
1,
)
g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)
return torch.nn.functional.grid_sample(
input=tenInput,
grid=g,
mode="bilinear",
padding_mode="border",
align_corners=True,
)
class BasicConv(nn.Module):
def __init__(
self,
in_planes,
out_planes,
kernel_size,
stride=1,
padding=0,
dilation=1,
groups=1,
relu=True,
bn=True,
bias=False,
):
super(BasicConv, self).__init__()
self.out_channels = out_planes
self.conv = nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups,
bias=bias,
)
self.bn = (
nn.BatchNorm2d(out_planes, eps=1e-5, momentum=0.01, affine=True)
if bn
else None
)
self.relu = nn.ReLU() if relu else None
def forward(self, x):
x = self.conv(x)
if self.bn is not None:
x = self.bn(x)
if self.relu is not None:
x = self.relu(x)
return x
class Flatten(nn.Module):
def forward(self, x):
return x.view(x.size(0), -1)
class ChannelGate(nn.Module):
def __init__(self, gate_channels, reduction_ratio=16, pool_types=["avg", "max"]):
super(ChannelGate, self).__init__()
self.gate_channels = gate_channels
self.mlp = nn.Sequential(
Flatten(),
nn.Linear(gate_channels, gate_channels // reduction_ratio),
nn.ReLU(),
nn.Linear(gate_channels // reduction_ratio, gate_channels),
)
self.pool_types = pool_types
def forward(self, x):
channel_att_sum = None
for pool_type in self.pool_types:
if pool_type == "avg":
avg_pool = F.avg_pool2d(
x, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))
)
channel_att_raw = self.mlp(avg_pool)
elif pool_type == "max":
max_pool = F.max_pool2d(
x, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))
)
channel_att_raw = self.mlp(max_pool)
elif pool_type == "lp":
lp_pool = F.lp_pool2d(
x, 2, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))
)
channel_att_raw = self.mlp(lp_pool)
elif pool_type == "lse":
# LSE pool only
lse_pool = logsumexp_2d(x)
channel_att_raw = self.mlp(lse_pool)
if channel_att_sum is None:
channel_att_sum = channel_att_raw
else:
channel_att_sum = channel_att_sum + channel_att_raw
scale = F.sigmoid(channel_att_sum).unsqueeze(2).unsqueeze(3).expand_as(x)
return x * scale
def logsumexp_2d(tensor):
tensor_flatten = tensor.view(tensor.size(0), tensor.size(1), -1)
s, _ = torch.max(tensor_flatten, dim=2, keepdim=True)
outputs = s + (tensor_flatten - s).exp().sum(dim=2, keepdim=True).log()
return outputs
class ChannelPool(nn.Module):
def forward(self, x):
return torch.cat(
(torch.max(x, 1)[0].unsqueeze(1), torch.mean(x, 1).unsqueeze(1)), dim=1
)
class SpatialGate(nn.Module):
def __init__(self):
super(SpatialGate, self).__init__()
kernel_size = 7
self.compress = ChannelPool()
self.spatial = BasicConv(
2, 1, kernel_size, stride=1, padding=(kernel_size - 1) // 2, relu=False
)
def forward(self, x):
x_compress = self.compress(x)
x_out = self.spatial(x_compress)
scale = F.sigmoid(x_out) # broadcasting
return x * scale
class CBAM(nn.Module):
def __init__(
self,
gate_channels,
reduction_ratio=16,
pool_types=["avg", "max"],
no_spatial=False,
):
super(CBAM, self).__init__()
self.ChannelGate = ChannelGate(gate_channels, reduction_ratio, pool_types)
self.no_spatial = no_spatial
if not no_spatial:
self.SpatialGate = SpatialGate()
def forward(self, x):
x_out = self.ChannelGate(x)
if not self.no_spatial:
x_out = self.SpatialGate(x_out)
return x_out
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
),
nn.PReLU(out_planes),
)
class UNetConv(nn.Module):
def __init__(self, in_planes, out_planes, att=True):
super(UNetConv, self).__init__()
self.conv1 = conv(in_planes, out_planes, 3, 2, 1)
self.conv2 = conv(out_planes, out_planes, 3, 1, 1)
if att:
self.cbam = CBAM(out_planes, 16) # 这一步导致了通道数最低为128
else:
self.cbam = None
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
if self.cbam is not None:
x = self.cbam(x)
return x
class UpConv(nn.Module):
def __init__(self, in_planes, out_planes, att=True):
super(UpConv, self).__init__()
self.deconv = nn.Sequential(
nn.ConvTranspose2d(in_planes, in_planes // 2, 4, 2, 1),
nn.PReLU(in_planes // 2),
)
# 也许不需要这么卷积,我不确定
self.conv1 = conv(in_planes, in_planes // 2, 3, 1, 1)
self.conv2 = conv(in_planes // 2, out_planes, 3, 1, 1)
if att:
self.cbam = CBAM(out_planes, 16)
else:
self.cbam = None
def forward(self, x1, x2):
x1 = self.deconv(x1)
y = self.conv1(torch.cat((x1, x2), 1))
y = self.conv2(y)
if self.cbam is not None:
y = self.cbam(y)
return y
class FeatureNet(nn.Module):
def __init__(self, in_planes, out_planes):
super(FeatureNet, self).__init__()
# 处理IFBlock0时通道数问题
self.conv0 = conv(7, in_planes, 1, 1, 0)
self.conv1 = UNetConv(in_planes, out_planes // 8, att=False)
self.conv2 = UNetConv(out_planes // 8, out_planes // 4, att=True)
self.conv3 = UNetConv(out_planes // 4, out_planes // 2, att=True)
self.conv4 = UNetConv(out_planes // 2, out_planes, att=True)
self.conv5 = UNetConv(out_planes, 2 * out_planes, att=True)
self.deconv5 = UpConv(2 * out_planes, out_planes, att=True)
self.deconv4 = UpConv(out_planes, out_planes // 2, att=False)
self.deconv3 = UpConv(out_planes // 2, out_planes // 4, att=False)
def forward(self, x, level=0):
if x.shape[1] != 17:
x = self.conv0(x)
x2 = self.conv1(x)
x4 = self.conv2(x2)
x8 = self.conv3(x4)
x16 = self.conv4(x8)
x32 = self.conv5(x16)
y = self.deconv5(x32, x16) # 匹配IFBlock0通道和尺寸
# “早退机制”以期待用同一个UNet提取特征,不确定是否对训练产生影响
if level != 0:
y = self.deconv4(y, x8) # 匹配IFBlock1通道和尺寸
if level == 2:
y = self.deconv3(y, x4) # 匹配IFBlock2通道和尺寸
return y
class IFBlock(nn.Module):
def __init__(self, c=64, level=0):
super(IFBlock, self).__init__()
self.convblock = nn.Sequential(
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
)
self.flowconv = nn.Conv2d(c, 4, 3, 1, 1)
self.maskconvx16 = nn.Conv2d(c, 16 * 16 * 9, 1, 1, 0)
self.maskconvx8 = nn.Conv2d(c, 8 * 8 * 9, 1, 1, 0)
self.maskconvx4 = nn.Conv2d(c, 4 * 4 * 9, 1, 1, 0)
self.level = level
assert self.level in [4, 8, 16], "Bitch"
def mask_conv(self, x):
if self.level == 4:
return self.maskconvx4(x)
if self.level == 8:
return self.maskconvx8(x)
if self.level == 16:
return self.maskconvx16(x)
def upsample_flow(self, flow, mask):
# 俺寻思俺懂了
N, _, H, W = flow.shape
mask = mask.view(N, 1, 9, self.level, self.level, H, W)
mask = torch.softmax(mask, dim=2)
up_flow = F.unfold(self.level * flow, [3, 3], padding=1)
up_flow = up_flow.view(N, 4, 9, 1, 1, H, W)
up_flow = torch.sum(mask * up_flow, dim=2)
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
return up_flow.reshape(N, 4, self.level * H, self.level * W)
def forward(self, x, scale):
x = self.convblock(x) + x # 类似ResNet的f(x) + x
tmp = self.flowconv(x)
up_mask = self.mask_conv(x)
flow_up = self.upsample_flow(tmp, up_mask)
flow = (
F.interpolate(
flow_up, scale_factor=scale, mode="bilinear", align_corners=False
)
* scale
)
return flow
class IFUNet(nn.Module):
def __init__(self):
super(IFUNet, self).__init__()
# block0通道数必须为128的整倍数
self.fmap = FeatureNet(in_planes=17, out_planes=256)
self.block0 = IFBlock(c=256, level=16)
self.block1 = IFBlock(c=128, level=8)
self.block2 = IFBlock(c=64, level=4)
def forward(self, x, scale=1.0, timestep=0.5, ensemble=True):
channel = x.shape[1] // 2
img0 = x[:, :channel]
img1 = x[:, channel:]
if not torch.is_tensor(timestep):
timestep = (x[:, :1].clone() * 0 + 1) * timestep
else:
timestep = timestep.repeat(1, 1, img0.shape[2], img0.shape[3])
warped_img0 = img0
warped_img1 = img1
flow = None
block = [self.block0, self.block1, self.block2]
for i in range(3):
if flow != None:
x = torch.cat((img0, img1, timestep, warped_img0, warped_img1), 1)
flowtmp = flow
if scale != 1:
x = F.interpolate(
x, scale_factor=scale, mode="bilinear", align_corners=False
)
flowtmp = (
F.interpolate(
flow,
scale_factor=scale,
mode="bilinear",
align_corners=False,
)
* scale
)
x = torch.cat((x, flowtmp), 1)
# 期待UNet能提取到特征,不再需要ensemble
Fmap = self.fmap(x, level=i)
flow_d = block[i](Fmap, scale=1.0 / scale)
flow = flow + flow_d
if ensemble:
x = torch.cat(
(img1, img0, 1 - timestep, warped_img0, warped_img1), 1
)
flowtmp = flow
if scale != 1:
x = F.interpolate(
x, scale_factor=scale, mode="bilinear", align_corners=False
)
flowtmp = (
F.interpolate(
flow,
scale_factor=scale,
mode="bilinear",
align_corners=False,
)
* scale
)
x = torch.cat((x, flowtmp), 1)
# 期待UNet能提取到特征,不再需要ensemble
Fmap = self.fmap(x, level=i)
flow_d = block[i](Fmap, scale=1.0 / scale)
flow2 = flow + flow_d
flow = (flow + flow2) / 2
else:
x = torch.cat((img0, img1, timestep), 1)
if scale != 1:
x = F.interpolate(
x, scale_factor=scale, mode="bilinear", align_corners=False
)
Fmap = self.fmap(x, level=i)
flow = block[i](Fmap, scale=1.0 / scale)
if ensemble:
x = torch.cat((img1, img0, 1 - timestep), 1)
if scale != 1:
x = F.interpolate(
x, scale_factor=scale, mode="bilinear", align_corners=False
)
Fmap = self.fmap(x, level=i)
flow2 = block[i](Fmap, scale=1.0 / scale)
flow = (flow + flow2) / 2
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
return flow, warped_img0, warped_img1
class IFUNetModel(nn.Module):
def __init__(self, local_rank=-1):
super(IFUNetModel, self).__init__()
self.flownet = IFUNet()
self.fusionnet = RRDBNet()
self.refinenet = ResynNet()
def forward(self, img0, img1, timestep=0.5, scale=1.0, ensemble=False):
n, c, h, w = img0.shape
ph = ((h - 1) // 64 + 1) * 64
pw = ((w - 1) // 64 + 1) * 64
padding = (0, pw - w, 0, ph - h)
img0 = F.pad(img0, padding)
img1 = F.pad(img1, padding)
imgs = torch.cat((img0, img1), 1)
flow, warped_img0, warped_img1 = self.flownet(imgs, scale, timestep, ensemble)
mask = self.fusionnet(img0, img1, warped_img0, warped_img1, flow)
merged = warped_img0 * mask + warped_img1 * (1 - mask)
merged, _ = self.refinenet(imgs, deg=merged, scale=[4, 2, 1])
return merged[:, :, :h, :w]
@@ -1,59 +0,0 @@
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
from comfy.model_management import get_torch_device
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["IFUNet.pth"]
class IFUnet_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 2, "max": 1000}),
"scale_factor": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 100, "step": 0.1}),
"ensemble": ("BOOLEAN", {"default":True})
},
"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: typing.SupportsInt = 1,
multiplier: typing.SupportsInt = 2,
scale_factor: typing.SupportsFloat = 1.0,
ensemble: bool = True,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
from .IFUNet_arch import IFUNetModel
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
interpolation_model = IFUNetModel()
interpolation_model.load_state_dict(torch.load(model_path))
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model, scale_factor, ensemble):
return model(frame_0, frame_1, timestep=timestep, scale=scale_factor, ensemble=ensemble)
args = [interpolation_model, scale_factor, ensemble]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
File diff suppressed because it is too large Load Diff
@@ -1,60 +0,0 @@
import pathlib
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames
import typing
from comfy.model_management import get_torch_device
from vfi_utils import InterpolationStateList, generic_frame_loop
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["M2M.pth"]
class M2M_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"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: typing.SupportsInt = 1,
multiplier: typing.SupportsInt = 2,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
from .M2M_arch import M2M_PWC
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
interpolation_model = M2M_PWC()
interpolation_model.load_state_dict(torch.load(model_path))
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, int_timestep, model):
tenSteps = [
torch.FloatTensor([int_timestep] * len(frame_0)).view(len(frame_0), 1, 1, 1).to(get_torch_device())
]
return model(frame_0, frame_1, tenSteps)[0]
args = [interpolation_model]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
@@ -1,22 +0,0 @@
import torch.multiprocessing as mp
if mp.current_process().name == "MainProcess":
import yaml
import os
from pathlib import Path
config_path = Path(Path(__file__).parent.parent.parent.resolve(), "config.yaml")
if os.path.exists(config_path):
config = yaml.load(open(config_path, "r"), Loader=yaml.FullLoader)
ops_backend = config["ops_backend"]
else:
ops_backend = "taichi"
assert ops_backend in ["taichi", "cupy"]
if ops_backend == "taichi":
from .taichi_ops import softsplat, ModuleSoftsplat, FunctionSoftsplat, softsplat_func, costvol_func, sepconv_func, init, batch_edt, FunctionAdaCoF, ModuleCorrelation, FunctionCorrelation, _FunctionCorrelation
else:
from .cupy_ops import softsplat, ModuleSoftsplat, FunctionSoftsplat, softsplat_func, costvol_func, sepconv_func, init, batch_edt, FunctionAdaCoF, ModuleCorrelation, FunctionCorrelation, _FunctionCorrelation
@@ -1,11 +0,0 @@
from .costvol import *
from .sepconv import *
from .softsplat import *
from .adacof import *
from .correlation import *
from comfy.model_management import is_nvidia, get_torch_device_name, get_torch_device
def init():
if not is_nvidia():
raise NotImplementedError(f"CuPy ops backend only support CUDA device but found {get_torch_device_name(get_torch_device())} instead. Try Taichi ops backend by editing config.yaml")
return
@@ -1,491 +0,0 @@
import torch
from .utils import cuda_kernel, cuda_launch, cuda_int32
import math
kernel_AdaCoF_updateOutput = """
extern "C" __global__ void kernel_AdaCoF_updateOutput(
const int n,
const float* input,
const float* weight,
const float* offset_i,
const float* offset_j,
float* output
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float dblOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(output) / SIZE_2(output) / SIZE_1(output) ) % SIZE_0(output);
const int c = ( intIndex / SIZE_3(output) / SIZE_2(output) ) % SIZE_1(output);
const int i = ( intIndex / SIZE_3(output) ) % SIZE_2(output);
const int j = ( intIndex ) % SIZE_3(output);
for (int k = 0; k < F_SIZE; k += 1) {
for (int l = 0; l < F_SIZE; l += 1) {
float w = VALUE_4(weight, intSample, k*F_SIZE+l, i, j);
float alpha = VALUE_4(offset_i, intSample, k*F_SIZE+l, i, j);
float beta = VALUE_4(offset_j, intSample, k*F_SIZE+l, i, j);
int A = (int) alpha;
int B = (int) beta;
int i_k_A = i+k*DILATION+A;
if(i_k_A < 0)
i_k_A = 0;
if(i_k_A > SIZE_2(input) - 1)
i_k_A = SIZE_2(input) - 1;
int j_l_B = j+l*DILATION+B;
if(j_l_B < 0)
j_l_B = 0;
if(j_l_B > SIZE_3(input) - 1)
j_l_B = SIZE_3(input) - 1;
int i_k_A_1 = i+k*DILATION+A+1;
if(i_k_A_1 < 0)
i_k_A_1 = 0;
if(i_k_A_1 > SIZE_2(input) - 1)
i_k_A_1 = SIZE_2(input) - 1;
int j_l_B_1 = j+l*DILATION+B+1;
if(j_l_B_1 < 0)
j_l_B_1 = 0;
if(j_l_B_1 > SIZE_3(input) - 1)
j_l_B_1 = SIZE_3(input) - 1;
dblOutput += w * (
VALUE_4(input, intSample, c, i_k_A, j_l_B)*(1-(alpha-(float)A))*(1-(beta-(float)B)) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B)*(alpha-(float)A)*(1-(beta-(float)B)) +
VALUE_4(input, intSample, c, i_k_A, j_l_B_1)*(1-(alpha-(float)A))*(beta-(float)B) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B_1)*(alpha-(float)A)*(beta-(float)B)
);
}
}
output[intIndex] = dblOutput;
} }
"""
kernel_AdaCoF_updateGradWeight = """
extern "C" __global__ void kernel_AdaCoF_updateGradWeight(
const int n,
const float* gradLoss,
const float* input,
const float* offset_i,
const float* offset_j,
float* gradWeight
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradWeight) / SIZE_2(gradWeight) / SIZE_1(gradWeight) ) % SIZE_0(gradWeight);
const int intDepth = ( intIndex / SIZE_3(gradWeight) / SIZE_2(gradWeight) ) % SIZE_1(gradWeight);
const int i = ( intIndex / SIZE_3(gradWeight) ) % SIZE_2(gradWeight);
const int j = ( intIndex ) % SIZE_3(gradWeight);
int k = intDepth / F_SIZE;
int l = intDepth % F_SIZE;
for (int c = 0; c < 3; c++)
{
float delta = VALUE_4(gradLoss, intSample, c, i, j);
float alpha = VALUE_4(offset_i, intSample, k*F_SIZE+l, i, j);
float beta = VALUE_4(offset_j, intSample, k*F_SIZE+l, i, j);
int A = (int) alpha;
int B = (int) beta;
int i_k_A = i+k*DILATION+A;
if(i_k_A < 0)
i_k_A = 0;
if(i_k_A > SIZE_2(input) - 1)
i_k_A = SIZE_2(input) - 1;
int j_l_B = j+l*DILATION+B;
if(j_l_B < 0)
j_l_B = 0;
if(j_l_B > SIZE_3(input) - 1)
j_l_B = SIZE_3(input) - 1;
int i_k_A_1 = i+k*DILATION+A+1;
if(i_k_A_1 < 0)
i_k_A_1 = 0;
if(i_k_A_1 > SIZE_2(input) - 1)
i_k_A_1 = SIZE_2(input) - 1;
int j_l_B_1 = j+l*DILATION+B+1;
if(j_l_B_1 < 0)
j_l_B_1 = 0;
if(j_l_B_1 > SIZE_3(input) - 1)
j_l_B_1 = SIZE_3(input) - 1;
floatOutput += delta * (
VALUE_4(input, intSample, c, i_k_A, j_l_B)*(1-(alpha-(float)A))*(1-(beta-(float)B)) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B)*(alpha-(float)A)*(1-(beta-(float)B)) +
VALUE_4(input, intSample, c, i_k_A, j_l_B_1)*(1-(alpha-(float)A))*(beta-(float)B) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B_1)*(alpha-(float)A)*(beta-(float)B)
);
}
gradWeight[intIndex] = floatOutput;
} }
"""
kernel_AdaCoF_updateGradAlpha = """
extern "C" __global__ void kernel_AdaCoF_updateGradAlpha(
const int n,
const float* gradLoss,
const float* input,
const float* weight,
const float* offset_i,
const float* offset_j,
float* gradOffset_i
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradOffset_i) / SIZE_2(gradOffset_i) / SIZE_1(gradOffset_i) ) % SIZE_0(gradOffset_i);
const int intDepth = ( intIndex / SIZE_3(gradOffset_i) / SIZE_2(gradOffset_i) ) % SIZE_1(gradOffset_i);
const int i = ( intIndex / SIZE_3(gradOffset_i) ) % SIZE_2(gradOffset_i);
const int j = ( intIndex ) % SIZE_3(gradOffset_i);
int k = intDepth / F_SIZE;
int l = intDepth % F_SIZE;
for (int c = 0; c < 3; c++)
{
float delta = VALUE_4(gradLoss, intSample, c, i, j);
float w = VALUE_4(weight, intSample, k*F_SIZE+l, i, j);
float alpha = VALUE_4(offset_i, intSample, k*F_SIZE+l, i, j);
float beta = VALUE_4(offset_j, intSample, k*F_SIZE+l, i, j);
int A = (int) alpha;
int B = (int) beta;
int i_k_A = i+k*DILATION+A;
if(i_k_A < 0)
i_k_A = 0;
if(i_k_A > SIZE_2(input) - 1)
i_k_A = SIZE_2(input) - 1;
int j_l_B = j+l*DILATION+B;
if(j_l_B < 0)
j_l_B = 0;
if(j_l_B > SIZE_3(input) - 1)
j_l_B = SIZE_3(input) - 1;
int i_k_A_1 = i+k*DILATION+A+1;
if(i_k_A_1 < 0)
i_k_A_1 = 0;
if(i_k_A_1 > SIZE_2(input) - 1)
i_k_A_1 = SIZE_2(input) - 1;
int j_l_B_1 = j+l*DILATION+B+1;
if(j_l_B_1 < 0)
j_l_B_1 = 0;
if(j_l_B_1 > SIZE_3(input) - 1)
j_l_B_1 = SIZE_3(input) - 1;
floatOutput += delta * w * (
- VALUE_4(input, intSample, c, i_k_A, j_l_B)*(1-(beta-(float)B)) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B)*(1-(beta-(float)B)) -
VALUE_4(input, intSample, c, i_k_A, j_l_B_1)*(beta-(float)B) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B_1)*(beta-(float)B)
);
}
gradOffset_i[intIndex] = floatOutput;
} }
"""
kernel_AdaCoF_updateGradBeta = """
extern "C" __global__ void kernel_AdaCoF_updateGradBeta(
const int n,
const float* gradLoss,
const float* input,
const float* weight,
const float* offset_i,
const float* offset_j,
float* gradOffset_j
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradOffset_j) / SIZE_2(gradOffset_j) / SIZE_1(gradOffset_j) ) % SIZE_0(gradOffset_j);
const int intDepth = ( intIndex / SIZE_3(gradOffset_j) / SIZE_2(gradOffset_j) ) % SIZE_1(gradOffset_j);
const int i = ( intIndex / SIZE_3(gradOffset_j) ) % SIZE_2(gradOffset_j);
const int j = ( intIndex ) % SIZE_3(gradOffset_j);
int k = intDepth / F_SIZE;
int l = intDepth % F_SIZE;
for (int c = 0; c < 3; c++)
{
float delta = VALUE_4(gradLoss, intSample, c, i, j);
float w = VALUE_4(weight, intSample, k*F_SIZE+l, i, j);
float alpha = VALUE_4(offset_i, intSample, k*F_SIZE+l, i, j);
float beta = VALUE_4(offset_j, intSample, k*F_SIZE+l, i, j);
int A = (int) alpha;
int B = (int) beta;
int i_k_A = i+k*DILATION+A;
if(i_k_A < 0)
i_k_A = 0;
if(i_k_A > SIZE_2(input) - 1)
i_k_A = SIZE_2(input) - 1;
int j_l_B = j+l*DILATION+B;
if(j_l_B < 0)
j_l_B = 0;
if(j_l_B > SIZE_3(input) - 1)
j_l_B = SIZE_3(input) - 1;
int i_k_A_1 = i+k*DILATION+A+1;
if(i_k_A_1 < 0)
i_k_A_1 = 0;
if(i_k_A_1 > SIZE_2(input) - 1)
i_k_A_1 = SIZE_2(input) - 1;
int j_l_B_1 = j+l*DILATION+B+1;
if(j_l_B_1 < 0)
j_l_B_1 = 0;
if(j_l_B_1 > SIZE_3(input) - 1)
j_l_B_1 = SIZE_3(input) - 1;
floatOutput += delta * w * (
- VALUE_4(input, intSample, c, i_k_A, j_l_B)*(1-(alpha-(float)A)) -
VALUE_4(input, intSample, c, i_k_A_1, j_l_B)*(alpha-(float)A) +
VALUE_4(input, intSample, c, i_k_A, j_l_B_1)*(1-(alpha-(float)A)) +
VALUE_4(input, intSample, c, i_k_A_1, j_l_B_1)*(alpha-(float)A)
);
}
gradOffset_j[intIndex] = floatOutput;
} }
"""
class FunctionAdaCoF(torch.autograd.Function):
# end
@staticmethod
def forward(ctx, input, weight, offset_i, offset_j, dilation):
ctx.save_for_backward(input, weight, offset_i, offset_j)
ctx.dilation = dilation
intSample = input.size(0)
intInputDepth = input.size(1)
intInputHeight = input.size(2)
intInputWidth = input.size(3)
intFilterSize = int(math.sqrt(weight.size(1)))
intOutputHeight = weight.size(2)
intOutputWidth = weight.size(3)
assert (
intInputHeight - ((intFilterSize - 1) * dilation + 1) == intOutputHeight - 1
)
assert (
intInputWidth - ((intFilterSize - 1) * dilation + 1) == intOutputWidth - 1
)
assert input.is_contiguous() == True
assert weight.is_contiguous() == True
assert offset_i.is_contiguous() == True
assert offset_j.is_contiguous() == True
output = input.new_zeros(
intSample, intInputDepth, intOutputHeight, intOutputWidth
)
if input.is_cuda == True:
class Stream:
ptr = torch.cuda.current_stream().cuda_stream
# end
n = output.nelement()
cuda_launch(
cuda_kernel(
"kernel_AdaCoF_updateOutput",
kernel_AdaCoF_updateOutput,
{
"input": input,
"weight": weight,
"offset_i": offset_i,
"offset_j": offset_j,
"output": output,
},
F_SIZE=str(intFilterSize),
DILATION=str(dilation)
),
)(
grid=tuple([int((n + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n,
input.data_ptr(),
weight.data_ptr(),
offset_i.data_ptr(),
offset_j.data_ptr(),
output.data_ptr(),
],
stream=Stream,
)
elif input.is_cuda == False:
raise NotImplementedError()
# end
return output
# end
@staticmethod
def backward(ctx, gradOutput):
input, weight, offset_i, offset_j = ctx.saved_tensors
dilation = ctx.dilation
intSample = input.size(0)
intInputDepth = input.size(1)
intInputHeight = input.size(2)
intInputWidth = input.size(3)
intFilterSize = int(math.sqrt(weight.size(1)))
intOutputHeight = weight.size(2)
intOutputWidth = weight.size(3)
assert (
intInputHeight - ((intFilterSize - 1) * dilation + 1) == intOutputHeight - 1
)
assert (
intInputWidth - ((intFilterSize - 1) * dilation + 1) == intOutputWidth - 1
)
assert gradOutput.is_contiguous() == True
gradInput = (
input.new_zeros(intSample, intInputDepth, intInputHeight, intInputWidth)
if ctx.needs_input_grad[0] == True
else None
)
gradWeight = (
input.new_zeros(
intSample, intFilterSize**2, intOutputHeight, intOutputWidth
)
if ctx.needs_input_grad[1] == True
else None
)
gradOffset_i = (
input.new_zeros(
intSample, intFilterSize**2, intOutputHeight, intOutputWidth
)
if ctx.needs_input_grad[2] == True
else None
)
gradOffset_j = (
input.new_zeros(
intSample, intFilterSize**2, intOutputHeight, intOutputWidth
)
if ctx.needs_input_grad[2] == True
else None
)
if input.is_cuda == True:
class Stream:
ptr = torch.cuda.current_stream().cuda_stream
# end
# weight grad
n_w = gradWeight.nelement()
cuda_launch(
cuda_kernel(
"kernel_AdaCoF_updateGradWeight",
kernel_AdaCoF_updateGradWeight,
{
"gradLoss": gradOutput,
"input": input,
"offset_i": offset_i,
"offset_j": offset_j,
"gradWeight": gradWeight,
},
F_SIZE=str(intFilterSize),
DILATION=str(dilation)
),
)(
grid=tuple([int((n_w + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n_w,
gradOutput.data_ptr(),
input.data_ptr(),
offset_i.data_ptr(),
offset_j.data_ptr(),
gradWeight.data_ptr(),
],
stream=Stream,
)
# alpha grad
n_i = gradOffset_i.nelement()
cuda_launch(
cuda_kernel(
"kernel_AdaCoF_updateGradAlpha",
kernel_AdaCoF_updateGradAlpha,
{
"gradLoss": gradOutput,
"input": input,
"weight": weight,
"offset_i": offset_i,
"offset_j": offset_j,
"gradOffset_i": gradOffset_i,
},
F_SIZE=str(intFilterSize),
DILATION=str(dilation)
),
)(
grid=tuple([int((n_i + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n_i,
gradOutput.data_ptr(),
input.data_ptr(),
weight.data_ptr(),
offset_i.data_ptr(),
offset_j.data_ptr(),
gradOffset_i.data_ptr(),
],
stream=Stream,
)
# beta grad
n_j = gradOffset_j.nelement()
cuda_launch(
cuda_kernel(
"kernel_AdaCoF_updateGradBeta",
kernel_AdaCoF_updateGradBeta,
{
"gradLoss": gradOutput,
"input": input,
"weight": weight,
"offset_i": offset_i,
"offset_j": offset_j,
"gradOffset_j": gradOffset_j,
},
F_SIZE=str(intFilterSize),
DILATION=str(dilation)
),
)(
grid=tuple([int((n_j + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n_j,
gradOutput.data_ptr(),
input.data_ptr(),
weight.data_ptr(),
offset_i.data_ptr(),
offset_j.data_ptr(),
gradOffset_j.data_ptr(),
],
stream=Stream,
)
elif input.is_cuda == False:
raise NotImplementedError()
# end
return gradInput, gradWeight, gradOffset_i, gradOffset_j, None
__all__ = ["FunctionAdaCoF"]
@@ -1,119 +0,0 @@
############### DISTANCE TRANSFORM ###############
# img tensor: (bs,h,w) or (bs,1,h,w)
# returns same shape
# expects white lines, black whitespace
# defaults to diameter if empty image
from .utils import cuda_kernel, cuda_launch, cuda_int32, cuda_float32
import torch
_batch_edt_kernel = (
"kernel_dt",
"""
extern "C" __global__ void kernel_dt(
const int bs,
const int h,
const int w,
const float diam2,
float* data,
float* output
) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= bs*h*w) {
return;
}
int pb = idx / (h*w);
int pi = (idx - h*w*pb) / w;
int pj = (idx - h*w*pb - w*pi);
float cost;
float mincost = diam2;
for (int j = 0; j < w; j++) {
cost = data[h*w*pb + w*pi + j] + (pj-j)*(pj-j);
if (cost < mincost) {
mincost = cost;
}
}
output[idx] = mincost;
return;
}
""",
)
_batch_edt = None
def batch_edt(img, block=1024):
# must initialize cuda/cupy after forking
global _batch_edt
if _batch_edt is None:
_batch_edt = cuda_launch(*_batch_edt_kernel)
# bookkeeppingg
if len(img.shape) == 4:
assert img.shape[1] == 1
img = img.squeeze(1)
expand = True
else:
expand = False
bs, h, w = img.shape
diam2 = h**2 + w**2
odtype = img.dtype
grid = (img.nelement() + block - 1) // block
# cupy implementation
if img.is_cuda:
# first pass, y-axis
data = ((1 - img.type(torch.float32)) * diam2).contiguous()
intermed = torch.zeros_like(data)
_batch_edt(
grid=(grid, 1, 1),
block=(block, 1, 1), # < 1024
args=[
cuda_int32(bs),
cuda_int32(h),
cuda_int32(w),
cuda_float32(diam2),
data.data_ptr(),
intermed.data_ptr(),
],
)
# second pass, x-axis
intermed = intermed.permute(0, 2, 1).contiguous()
out = torch.zeros_like(intermed)
_batch_edt(
grid=(grid, 1, 1),
block=(block, 1, 1),
args=[
cuda_int32(bs),
cuda_int32(w),
cuda_int32(h),
cuda_float32(diam2),
intermed.data_ptr(),
out.data_ptr(),
],
)
ans = out.permute(0, 2, 1).sqrt()
ans = ans.type(odtype) if odtype != ans.dtype else ans
# default to scipy cpu implementation
else:
raise NotImplementedError()
""" sums = img.sum(dim=(1, 2))
ans = torch.tensor(
np.stack(
[
scipy.ndimage.morphology.distance_transform_edt(i)
if s != 0
else np.ones_like(i) # change scipy behavior for empty image
* np.sqrt(diam2)
for i, s in zip(1 - img, sums)
]
),
dtype=odtype,
) """
if expand:
ans = ans.unsqueeze(1)
return ans
__all__ = ["batch_edt"]
@@ -1,413 +0,0 @@
import torch
from .utils import cuda_kernel, cuda_launch, cuda_int32
kernel_Correlation_rearrange = """
extern "C" __global__ void kernel_Correlation_rearrange(
const int n,
const float* input,
float* output
) {
int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x;
if (intIndex >= n) {
return;
}
int intSample = blockIdx.z;
int intChannel = blockIdx.y;
float fltValue = input[(((intSample * SIZE_1(input)) + intChannel) * SIZE_2(input) * SIZE_3(input)) + intIndex];
__syncthreads();
int intPaddedY = (intIndex / SIZE_3(input)) + 4;
int intPaddedX = (intIndex % SIZE_3(input)) + 4;
int intRearrange = ((SIZE_3(input) + 8) * intPaddedY) + intPaddedX;
output[(((intSample * SIZE_1(output) * SIZE_2(output)) + intRearrange) * SIZE_1(input)) + intChannel] = fltValue;
}
"""
kernel_Correlation_updateOutput = """
extern "C" __global__ void kernel_Correlation_updateOutput(
const int n,
const float* rbot0,
const float* rbot1,
float* top
) {
extern __shared__ char patch_data_char[];
float *patch_data = (float *)patch_data_char;
// First (upper left) position of kernel upper-left corner in current center position of neighborhood in image 1
int x1 = blockIdx.x + 4;
int y1 = blockIdx.y + 4;
int item = blockIdx.z;
int ch_off = threadIdx.x;
// Load 3D patch into shared shared memory
for (int j = 0; j < 1; j++) { // HEIGHT
for (int i = 0; i < 1; i++) { // WIDTH
int ji_off = (j + i) * SIZE_3(rbot0);
for (int ch = ch_off; ch < SIZE_3(rbot0); ch += 32) { // CHANNELS
int idx1 = ((item * SIZE_1(rbot0) + y1+j) * SIZE_2(rbot0) + x1+i) * SIZE_3(rbot0) + ch;
int idxPatchData = ji_off + ch;
patch_data[idxPatchData] = rbot0[idx1];
}
}
}
__syncthreads();
__shared__ float sum[32];
// Compute correlation
for (int top_channel = 0; top_channel < SIZE_1(top); top_channel++) {
sum[ch_off] = 0;
int s2o = top_channel % 9 - 4;
int s2p = top_channel / 9 - 4;
for (int j = 0; j < 1; j++) { // HEIGHT
for (int i = 0; i < 1; i++) { // WIDTH
int ji_off = (j + i) * SIZE_3(rbot0);
for (int ch = ch_off; ch < SIZE_3(rbot0); ch += 32) { // CHANNELS
int x2 = x1 + s2o;
int y2 = y1 + s2p;
int idxPatchData = ji_off + ch;
int idx2 = ((item * SIZE_1(rbot0) + y2+j) * SIZE_2(rbot0) + x2+i) * SIZE_3(rbot0) + ch;
sum[ch_off] += patch_data[idxPatchData] * rbot1[idx2];
}
}
}
__syncthreads();
if (ch_off == 0) {
float total_sum = 0;
for (int idx = 0; idx < 32; idx++) {
total_sum += sum[idx];
}
const int sumelems = SIZE_3(rbot0);
const int index = ((top_channel*SIZE_2(top) + blockIdx.y)*SIZE_3(top))+blockIdx.x;
top[index + item*SIZE_1(top)*SIZE_2(top)*SIZE_3(top)] = total_sum / (float)sumelems;
}
}
}
"""
kernel_Correlation_updateGradFirst = """
#define ROUND_OFF 50000
extern "C" __global__ void kernel_Correlation_updateGradFirst(
const int n,
const int intSample,
const float* rbot0,
const float* rbot1,
const float* gradOutput,
float* gradFirst,
float* gradSecond
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
int n = intIndex % SIZE_1(gradFirst); // channels
int l = (intIndex / SIZE_1(gradFirst)) % SIZE_3(gradFirst) + 4; // w-pos
int m = (intIndex / SIZE_1(gradFirst) / SIZE_3(gradFirst)) % SIZE_2(gradFirst) + 4; // h-pos
// round_off is a trick to enable integer division with ceil, even for negative numbers
// We use a large offset, for the inner part not to become negative.
const int round_off = ROUND_OFF;
const int round_off_s1 = round_off;
// We add round_off before_s1 the int division and subtract round_off after it, to ensure the formula matches ceil behavior:
int xmin = (l - 4 + round_off_s1 - 1) + 1 - round_off; // ceil (l - 4)
int ymin = (m - 4 + round_off_s1 - 1) + 1 - round_off; // ceil (l - 4)
// Same here:
int xmax = (l - 4 + round_off_s1) - round_off; // floor (l - 4)
int ymax = (m - 4 + round_off_s1) - round_off; // floor (m - 4)
float sum = 0;
if (xmax>=0 && ymax>=0 && (xmin<=SIZE_3(gradOutput)-1) && (ymin<=SIZE_2(gradOutput)-1)) {
xmin = max(0,xmin);
xmax = min(SIZE_3(gradOutput)-1,xmax);
ymin = max(0,ymin);
ymax = min(SIZE_2(gradOutput)-1,ymax);
for (int p = -4; p <= 4; p++) {
for (int o = -4; o <= 4; o++) {
// Get rbot1 data:
int s2o = o;
int s2p = p;
int idxbot1 = ((intSample * SIZE_1(rbot0) + (m+s2p)) * SIZE_2(rbot0) + (l+s2o)) * SIZE_3(rbot0) + n;
float bot1tmp = rbot1[idxbot1]; // rbot1[l+s2o,m+s2p,n]
// Index offset for gradOutput in following loops:
int op = (p+4) * 9 + (o+4); // index[o,p]
int idxopoffset = (intSample * SIZE_1(gradOutput) + op);
for (int y = ymin; y <= ymax; y++) {
for (int x = xmin; x <= xmax; x++) {
int idxgradOutput = (idxopoffset * SIZE_2(gradOutput) + y) * SIZE_3(gradOutput) + x; // gradOutput[x,y,o,p]
sum += gradOutput[idxgradOutput] * bot1tmp;
}
}
}
}
}
const int sumelems = SIZE_1(gradFirst);
const int bot0index = ((n * SIZE_2(gradFirst)) + (m-4)) * SIZE_3(gradFirst) + (l-4);
gradFirst[bot0index + intSample*SIZE_1(gradFirst)*SIZE_2(gradFirst)*SIZE_3(gradFirst)] = sum / (float)sumelems;
} }
"""
kernel_Correlation_updateGradSecond = """
#define ROUND_OFF 50000
extern "C" __global__ void kernel_Correlation_updateGradSecond(
const int n,
const int intSample,
const float* rbot0,
const float* rbot1,
const float* gradOutput,
float* gradFirst,
float* gradSecond
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
int n = intIndex % SIZE_1(gradSecond); // channels
int l = (intIndex / SIZE_1(gradSecond)) % SIZE_3(gradSecond) + 4; // w-pos
int m = (intIndex / SIZE_1(gradSecond) / SIZE_3(gradSecond)) % SIZE_2(gradSecond) + 4; // h-pos
// round_off is a trick to enable integer division with ceil, even for negative numbers
// We use a large offset, for the inner part not to become negative.
const int round_off = ROUND_OFF;
const int round_off_s1 = round_off;
float sum = 0;
for (int p = -4; p <= 4; p++) {
for (int o = -4; o <= 4; o++) {
int s2o = o;
int s2p = p;
//Get X,Y ranges and clamp
// We add round_off before_s1 the int division and subtract round_off after it, to ensure the formula matches ceil behavior:
int xmin = (l - 4 - s2o + round_off_s1 - 1) + 1 - round_off; // ceil (l - 4 - s2o)
int ymin = (m - 4 - s2p + round_off_s1 - 1) + 1 - round_off; // ceil (l - 4 - s2o)
// Same here:
int xmax = (l - 4 - s2o + round_off_s1) - round_off; // floor (l - 4 - s2o)
int ymax = (m - 4 - s2p + round_off_s1) - round_off; // floor (m - 4 - s2p)
if (xmax>=0 && ymax>=0 && (xmin<=SIZE_3(gradOutput)-1) && (ymin<=SIZE_2(gradOutput)-1)) {
xmin = max(0,xmin);
xmax = min(SIZE_3(gradOutput)-1,xmax);
ymin = max(0,ymin);
ymax = min(SIZE_2(gradOutput)-1,ymax);
// Get rbot0 data:
int idxbot0 = ((intSample * SIZE_1(rbot0) + (m-s2p)) * SIZE_2(rbot0) + (l-s2o)) * SIZE_3(rbot0) + n;
float bot0tmp = rbot0[idxbot0]; // rbot1[l+s2o,m+s2p,n]
// Index offset for gradOutput in following loops:
int op = (p+4) * 9 + (o+4); // index[o,p]
int idxopoffset = (intSample * SIZE_1(gradOutput) + op);
for (int y = ymin; y <= ymax; y++) {
for (int x = xmin; x <= xmax; x++) {
int idxgradOutput = (idxopoffset * SIZE_2(gradOutput) + y) * SIZE_3(gradOutput) + x; // gradOutput[x,y,o,p]
sum += gradOutput[idxgradOutput] * bot0tmp;
}
}
}
}
}
const int sumelems = SIZE_1(gradSecond);
const int bot1index = ((n * SIZE_2(gradSecond)) + (m-4)) * SIZE_3(gradSecond) + (l-4);
gradSecond[bot1index + intSample*SIZE_1(gradSecond)*SIZE_2(gradSecond)*SIZE_3(gradSecond)] = sum / (float)sumelems;
} }
"""
class _FunctionCorrelation(torch.autograd.Function):
@staticmethod
def forward(self, first, second):
rbot0 = first.new_zeros(
[first.shape[0], first.shape[2] + 8, first.shape[3] + 8, first.shape[1]]
)
rbot1 = first.new_zeros(
[first.shape[0], first.shape[2] + 8, first.shape[3] + 8, first.shape[1]]
)
self.save_for_backward(first, second, rbot0, rbot1)
first = first.contiguous()
assert first.is_cuda == True
second = second.contiguous()
assert second.is_cuda == True
output = first.new_zeros([first.shape[0], 81, first.shape[2], first.shape[3]])
if first.is_cuda == True:
n = first.shape[2] * first.shape[3]
cuda_launch(
cuda_kernel(
"kernel_Correlation_rearrange", kernel_Correlation_rearrange, {"input": first, "output": rbot0}
),
)(
grid=tuple([int((n + 16 - 1) / 16), first.shape[1], first.shape[0]]),
block=tuple([16, 1, 1]),
args=[n, first.data_ptr(), rbot0.data_ptr()],
)
n = second.shape[2] * second.shape[3]
cuda_launch(
cuda_kernel(
"kernel_Correlation_rearrange", kernel_Correlation_rearrange, {"input": second, "output": rbot1}
),
)(
grid=tuple([int((n + 16 - 1) / 16), second.shape[1], second.shape[0]]),
block=tuple([16, 1, 1]),
args=[n, second.data_ptr(), rbot1.data_ptr()],
)
n = output.shape[1] * output.shape[2] * output.shape[3]
cuda_launch(
cuda_kernel(
"kernel_Correlation_updateOutput",
kernel_Correlation_updateOutput,
{"rbot0": rbot0, "rbot1": rbot1, "top": output},
),
)(
grid=tuple([output.shape[3], output.shape[2], output.shape[0]]),
block=tuple([32, 1, 1]),
shared_mem=first.shape[1] * 4,
args=[n, rbot0.data_ptr(), rbot1.data_ptr(), output.data_ptr()],
)
elif first.is_cuda == False:
raise NotImplementedError()
# end
return output
# end
@staticmethod
def backward(self, gradOutput):
first, second, rbot0, rbot1 = self.saved_tensors
gradOutput = gradOutput.contiguous()
assert gradOutput.is_cuda == True
gradFirst = (
first.new_zeros(
[first.shape[0], first.shape[1], first.shape[2], first.shape[3]]
)
if self.needs_input_grad[0] == True
else None
)
gradSecond = (
first.new_zeros(
[first.shape[0], first.shape[1], first.shape[2], first.shape[3]]
)
if self.needs_input_grad[1] == True
else None
)
if first.is_cuda == True:
if gradFirst is not None:
for intSample in range(first.shape[0]):
n = first.shape[1] * first.shape[2] * first.shape[3]
cuda_launch(
cuda_kernel(
"kernel_Correlation_updateGradFirst",
kernel_Correlation_updateGradFirst,
{
"rbot0": rbot0,
"rbot1": rbot1,
"gradOutput": gradOutput,
"gradFirst": gradFirst,
"gradSecond": None,
},
),
)(
grid=tuple([int((n + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n,
intSample,
rbot0.data_ptr(),
rbot1.data_ptr(),
gradOutput.data_ptr(),
gradFirst.data_ptr(),
None,
],
)
# end
# end
if gradSecond is not None:
for intSample in range(first.shape[0]):
n = first.shape[1] * first.shape[2] * first.shape[3]
cuda_launch(
cuda_kernel(
"kernel_Correlation_updateGradSecond",
kernel_Correlation_updateGradSecond,
{
"rbot0": rbot0,
"rbot1": rbot1,
"gradOutput": gradOutput,
"gradFirst": None,
"gradSecond": gradSecond,
},
),
)(
grid=tuple([int((n + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
n,
intSample,
rbot0.data_ptr(),
rbot1.data_ptr(),
gradOutput.data_ptr(),
None,
gradSecond.data_ptr(),
],
)
# end
# end
elif first.is_cuda == False:
raise NotImplementedError()
# end
return gradFirst, gradSecond
# end
# end
def FunctionCorrelation(tenFirst, tenSecond):
return _FunctionCorrelation.apply(tenFirst, tenSecond)
# end
class ModuleCorrelation(torch.nn.Module):
def __init__(self):
super(ModuleCorrelation, self).__init__()
# end
def forward(self, tenFirst, tenSecond):
return _FunctionCorrelation.apply(tenFirst, tenSecond)
# end
__all__ = ["_FunctionCorrelation", "FunctionCorrelation", "ModuleCorrelation"]
@@ -1,317 +0,0 @@
from .utils import cuda_kernel, cuda_launch, cuda_int32
import torch, collections
costvol_out = """
extern "C" __global__ void __launch_bounds__(512) costvol_out(
const int n,
const {{type}}* __restrict__ tenOne,
const {{type}}* __restrict__ tenTwo,
{{type}}* __restrict__ tenOut
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) ) % SIZE_0(tenOut);
const int intC = -1;
const int intY = ( intIndex / SIZE_3(tenOut) ) % SIZE_2(tenOut);
const int intX = ( intIndex ) % SIZE_3(tenOut);
{{type}} fltOne[{{intChans}}];
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
fltOne[intValue] = VALUE_4(tenOne, intN, intValue, intY, intX);
}
int intOffset = OFFSET_4(tenOut, intN, 0, intY, intX);
for (int intOy = intY - 4; intOy <= intY + 4; intOy += 1) {
for (int intOx = intX - 4; intOx <= intX + 4; intOx += 1) {
{{type}} fltValue = 0.0f;
if ((intOy >= 0) && (intOy < SIZE_2(tenOut)) && (intOx >= 0) && (intOx < SIZE_3(tenOut))) {
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
fltValue += abs(fltOne[intValue] - VALUE_4(tenTwo, intN, intValue, intOy, intOx));
}
} else {
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
fltValue += abs(fltOne[intValue]);
}
}
tenOut[intOffset] = fltValue / SIZE_1(tenOne);
intOffset += SIZE_2(tenOut) * SIZE_3(tenOut);
}
}
} }
"""
costvol_onegrad = """
extern "C" __global__ void __launch_bounds__(512) costvol_onegrad(
const int n,
const {{type}}* __restrict__ tenOne,
const {{type}}* __restrict__ tenTwo,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenOnegrad,
{{type}}* __restrict__ tenTwograd
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenOnegrad) / SIZE_2(tenOnegrad) ) % SIZE_0(tenOnegrad);
const int intC = -1;
const int intY = ( intIndex / SIZE_3(tenOnegrad) ) % SIZE_2(tenOnegrad);
const int intX = ( intIndex ) % SIZE_3(tenOnegrad);
{{type}} fltOne[{{intChans}}];
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
fltOne[intValue] = VALUE_4(tenOne, intN, intValue, intY, intX);
}
int intOffset = OFFSET_4(tenOutgrad, intN, 0, intY, intX);
for (int intOy = intY - 4; intOy <= intY + 4; intOy += 1) {
for (int intOx = intX - 4; intOx <= intX + 4; intOx += 1) {
if ((intOy >= 0) && (intOy < SIZE_2(tenOutgrad)) && (intOx >= 0) && (intOx < SIZE_3(tenOutgrad))) {
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
if (fltOne[intValue] - VALUE_4(tenTwo, intN, intValue, intOy, intOx) >= 0.0f) {
tenOnegrad[OFFSET_4(tenOnegrad, intN, intValue, intY, intX)] += +tenOutgrad[intOffset] / SIZE_1(tenOne);
} else {
tenOnegrad[OFFSET_4(tenOnegrad, intN, intValue, intY, intX)] += -tenOutgrad[intOffset] / SIZE_1(tenOne);
}
}
} else {
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
if (fltOne[intValue] >= 0.0f) {
tenOnegrad[OFFSET_4(tenOnegrad, intN, intValue, intY, intX)] += +tenOutgrad[intOffset] / SIZE_1(tenOne);
} else {
tenOnegrad[OFFSET_4(tenOnegrad, intN, intValue, intY, intX)] += -tenOutgrad[intOffset] / SIZE_1(tenOne);
}
}
}
intOffset += SIZE_2(tenOutgrad) * SIZE_3(tenOutgrad);
}
}
} }
"""
costvol_twograd = """
extern "C" __global__ void __launch_bounds__(512) costvol_twograd(
const int n,
const {{type}}* __restrict__ tenOne,
const {{type}}* __restrict__ tenTwo,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenOnegrad,
{{type}}* __restrict__ tenTwograd
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenTwograd) / SIZE_2(tenTwograd) ) % SIZE_0(tenTwograd);
const int intC = -1;
const int intY = ( intIndex / SIZE_3(tenTwograd) ) % SIZE_2(tenTwograd);
const int intX = ( intIndex ) % SIZE_3(tenTwograd);
{{type}} fltOne[{{intChans}}];
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
fltOne[intValue] = VALUE_4(tenOne, intN, intValue, intY, intX);
}
int intOffset = OFFSET_4(tenOutgrad, intN, 0, intY, intX);
for (int intOy = intY - 4; intOy <= intY + 4; intOy += 1) {
for (int intOx = intX - 4; intOx <= intX + 4; intOx += 1) {
if ((intOy >= 0) && (intOy < SIZE_2(tenOutgrad)) && (intOx >= 0) && (intOx < SIZE_3(tenOutgrad))) {
for (int intValue = 0; intValue < SIZE_1(tenOne); intValue += 1) {
if (fltOne[intValue] - VALUE_4(tenTwo, intN, intValue, intOy, intOx) >= 0.0f) {
atomicAdd(&tenTwograd[OFFSET_4(tenTwograd, intN, intValue, intOy, intOx)], -tenOutgrad[intOffset] / SIZE_1(tenOne));
} else {
atomicAdd(&tenTwograd[OFFSET_4(tenTwograd, intN, intValue, intOy, intOx)], +tenOutgrad[intOffset] / SIZE_1(tenOne));
}
}
} else {
// ...
}
intOffset += SIZE_2(tenOutgrad) * SIZE_3(tenOutgrad);
}
}
} }
"""
class costvol_func(torch.autograd.Function):
@staticmethod
@torch.cuda.amp.custom_fwd(cast_inputs=torch.float32)
def forward(self, tenOne, tenTwo):
tenOut = tenOne.new_empty(
[tenOne.shape[0], 81, tenOne.shape[2], tenOne.shape[3]]
)
cuda_launch(
cuda_kernel(
"costvol_out",
costvol_out,
{
"intChans": tenOne.shape[1],
"tenOne": tenOne,
"tenTwo": tenTwo,
"tenOut": tenOut,
},
)
)(
grid=tuple(
[
int(
(
(tenOut.shape[0] * tenOut.shape[2] * tenOut.shape[3])
+ 512
- 1
)
/ 512
),
1,
1,
]
),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenOut.shape[0] * tenOut.shape[2] * tenOut.shape[3]),
tenOne.data_ptr(),
tenTwo.data_ptr(),
tenOut.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
self.save_for_backward(tenOne, tenTwo)
return tenOut
# end
@staticmethod
@torch.cuda.amp.custom_bwd
def backward(self, tenOutgrad):
tenOne, tenTwo = self.saved_tensors
tenOutgrad = tenOutgrad.contiguous()
assert tenOutgrad.is_cuda == True
tenOnegrad = (
tenOne.new_zeros(
[tenOne.shape[0], tenOne.shape[1], tenOne.shape[2], tenOne.shape[3]]
)
if self.needs_input_grad[0] == True
else None
)
tenTwograd = (
tenTwo.new_zeros(
[tenTwo.shape[0], tenTwo.shape[1], tenTwo.shape[2], tenTwo.shape[3]]
)
if self.needs_input_grad[1] == True
else None
)
if tenOnegrad is not None:
cuda_launch(
cuda_kernel(
"costvol_onegrad",
costvol_onegrad,
{
"intChans": tenOne.shape[1],
"tenOne": tenOne,
"tenTwo": tenTwo,
"tenOutgrad": tenOutgrad,
"tenOnegrad": tenOnegrad,
"tenTwograd": tenTwograd,
},
)
)(
grid=tuple(
[
int(
(
(
tenOnegrad.shape[0]
* tenOnegrad.shape[2]
* tenOnegrad.shape[3]
)
+ 512
- 1
)
/ 512
),
1,
1,
]
),
block=tuple([512, 1, 1]),
args=[
cuda_int32(
tenOnegrad.shape[0] * tenOnegrad.shape[2] * tenOnegrad.shape[3]
),
tenOne.data_ptr(),
tenTwo.data_ptr(),
tenOutgrad.data_ptr(),
tenOnegrad.data_ptr(),
tenTwograd.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
if tenTwograd is not None:
cuda_launch(
cuda_kernel(
"costvol_twograd",
costvol_twograd,
{
"intChans": tenOne.shape[1],
"tenOne": tenOne,
"tenTwo": tenTwo,
"tenOutgrad": tenOutgrad,
"tenOnegrad": tenOnegrad,
"tenTwograd": tenTwograd,
},
)
)(
grid=tuple(
[
int(
(
(
tenTwograd.shape[0]
* tenTwograd.shape[2]
* tenTwograd.shape[3]
)
+ 512
- 1
)
/ 512
),
1,
1,
]
),
block=tuple([512, 1, 1]),
args=[
cuda_int32(
tenTwograd.shape[0] * tenTwograd.shape[2] * tenTwograd.shape[3]
),
tenOne.data_ptr(),
tenTwo.data_ptr(),
tenOutgrad.data_ptr(),
tenOnegrad.data_ptr(),
tenTwograd.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
return tenOnegrad, tenTwograd, None, None
# end
# end
__all__ = ["costvol_func"]
@@ -1,332 +0,0 @@
import torch
from .utils import cuda_launch, cuda_kernel, cuda_int32
sepconv_vergrad = """
extern "C" __global__ void __launch_bounds__(512) sepconv_vergrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenVer,
const {{type}}* __restrict__ tenHor,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenVergrad,
{{type}}* __restrict__ tenHorgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenVergrad) / SIZE_2(tenVergrad) / SIZE_1(tenVergrad) ) % SIZE_0(tenVergrad);
const int intC = ( intIndex / SIZE_3(tenVergrad) / SIZE_2(tenVergrad) ) % SIZE_1(tenVergrad);
const int intY = ( intIndex / SIZE_3(tenVergrad) ) % SIZE_2(tenVergrad);
const int intX = ( intIndex ) % SIZE_3(tenVergrad);
{{type}} fltVergrad = 0.0;
{{type}} fltKahanc = 0.0;
{{type}} fltKahany = 0.0;
{{type}} fltKahant = 0.0;
for (int intI = 0; intI < SIZE_1(tenIn); intI += 1) {
for (int intFx = 0; intFx < SIZE_1(tenHor); intFx += 1) {
fltKahany = VALUE_4(tenHor, intN, intFx, intY, intX) * VALUE_4(tenIn, intN, intI, intY + intC, intX + intFx) * VALUE_4(tenOutgrad, intN, intI, intY, intX);
fltKahany = fltKahany - fltKahanc;
fltKahant = fltVergrad + fltKahany;
fltKahanc = (fltKahant - fltVergrad) - fltKahany;
fltVergrad = fltKahant;
}
}
tenVergrad[intIndex] = fltVergrad;
} }
"""
sepconv_ingrad = """
extern "C" __global__ void __launch_bounds__(512) sepconv_ingrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenVer,
const {{type}}* __restrict__ tenHor,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenVergrad,
{{type}}* __restrict__ tenHorgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) / SIZE_1(tenIngrad) ) % SIZE_0(tenIngrad);
const int intC = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) ) % SIZE_1(tenIngrad);
const int intY = ( intIndex / SIZE_3(tenIngrad) ) % SIZE_2(tenIngrad);
const int intX = ( intIndex ) % SIZE_3(tenIngrad);
{{type}} fltIngrad = 0.0;
{{type}} fltKahanc = 0.0;
{{type}} fltKahany = 0.0;
{{type}} fltKahant = 0.0;
for (int intFy = 0; intFy < SIZE_1(tenVer); intFy += 1) {
int intKy = intY + intFy - (SIZE_1(tenVer) - 1);
if (intKy < 0) { continue; }
if (intKy >= SIZE_2(tenVer)) { continue; }
for (int intFx = 0; intFx < SIZE_1(tenHor); intFx += 1) {
int intKx = intX + intFx - (SIZE_1(tenHor) - 1);
if (intKx < 0) { continue; }
if (intKx >= SIZE_3(tenHor)) { continue; }
fltKahany = VALUE_4(tenVer, intN, (SIZE_1(tenVer) - 1) - intFy, intKy, intKx) * VALUE_4(tenHor, intN, (SIZE_1(tenHor) - 1) - intFx, intKy, intKx) * VALUE_4(tenOutgrad, intN, intC, intKy, intKx);
fltKahany = fltKahany - fltKahanc;
fltKahant = fltIngrad + fltKahany;
fltKahanc = (fltKahant - fltIngrad) - fltKahany;
fltIngrad = fltKahant;
}
}
tenIngrad[intIndex] = fltIngrad;
} }
"""
sepconv_out = """
extern "C" __global__ void __launch_bounds__(512) sepconv_out(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenVer,
const {{type}}* __restrict__ tenHor,
{{type}}* __restrict__ tenOut
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) / SIZE_1(tenOut) ) % SIZE_0(tenOut);
const int intC = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) ) % SIZE_1(tenOut);
const int intY = ( intIndex / SIZE_3(tenOut) ) % SIZE_2(tenOut);
const int intX = ( intIndex ) % SIZE_3(tenOut);
{{type}} fltOut = 0.0;
{{type}} fltKahanc = 0.0;
{{type}} fltKahany = 0.0;
{{type}} fltKahant = 0.0;
for (int intFy = 0; intFy < SIZE_1(tenVer); intFy += 1) {
for (int intFx = 0; intFx < SIZE_1(tenHor); intFx += 1) {
fltKahany = VALUE_4(tenIn, intN, intC, intY + intFy, intX + intFx) * VALUE_4(tenVer, intN, intFy, intY, intX) * VALUE_4(tenHor, intN, intFx, intY, intX);
fltKahany = fltKahany - fltKahanc;
fltKahant = fltOut + fltKahany;
fltKahanc = (fltKahant - fltOut) - fltKahany;
fltOut = fltKahant;
}
}
tenOut[intIndex] = fltOut;
} }
"""
sepconv_horgrad = """
extern "C" __global__ void __launch_bounds__(512) sepconv_horgrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenVer,
const {{type}}* __restrict__ tenHor,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenVergrad,
{{type}}* __restrict__ tenHorgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenHorgrad) / SIZE_2(tenHorgrad) / SIZE_1(tenHorgrad) ) % SIZE_0(tenHorgrad);
const int intC = ( intIndex / SIZE_3(tenHorgrad) / SIZE_2(tenHorgrad) ) % SIZE_1(tenHorgrad);
const int intY = ( intIndex / SIZE_3(tenHorgrad) ) % SIZE_2(tenHorgrad);
const int intX = ( intIndex ) % SIZE_3(tenHorgrad);
{{type}} fltHorgrad = 0.0;
{{type}} fltKahanc = 0.0;
{{type}} fltKahany = 0.0;
{{type}} fltKahant = 0.0;
for (int intI = 0; intI < SIZE_1(tenIn); intI += 1) {
for (int intFy = 0; intFy < SIZE_1(tenVer); intFy += 1) {
fltKahany = VALUE_4(tenVer, intN, intFy, intY, intX) * VALUE_4(tenIn, intN, intI, intY + intFy, intX + intC) * VALUE_4(tenOutgrad, intN, intI, intY, intX);
fltKahany = fltKahany - fltKahanc;
fltKahant = fltHorgrad + fltKahany;
fltKahanc = (fltKahant - fltHorgrad) - fltKahany;
fltHorgrad = fltKahant;
}
}
tenHorgrad[intIndex] = fltHorgrad;
} }
"""
class sepconv_func(torch.autograd.Function):
@staticmethod
@torch.cuda.amp.custom_fwd(cast_inputs=torch.float32)
def forward(self, tenIn, tenVer, tenHor):
tenOut = tenIn.new_empty(
[
tenIn.shape[0],
tenIn.shape[1],
tenVer.shape[2] and tenHor.shape[2],
tenVer.shape[3] and tenHor.shape[3],
]
)
if tenIn.is_cuda == True:
cuda_launch(
cuda_kernel(
"sepconv_out",
sepconv_out,
{
"tenIn": tenIn,
"tenVer": tenVer,
"tenHor": tenHor,
"tenOut": tenOut,
},
)
)(
grid=tuple([int((tenOut.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenOut.nelement()),
tenIn.data_ptr(),
tenVer.data_ptr(),
tenHor.data_ptr(),
tenOut.data_ptr(),
],
)
elif tenIn.is_cuda != True:
assert False
# end
self.save_for_backward(tenIn, tenVer, tenHor)
return tenOut
# end
@staticmethod
@torch.cuda.amp.custom_bwd
def backward(self, tenOutgrad):
tenIn, tenVer, tenHor = self.saved_tensors
tenOutgrad = tenOutgrad.contiguous()
assert tenOutgrad.is_cuda == True
tenIngrad = (
tenIn.new_empty(
[tenIn.shape[0], tenIn.shape[1], tenIn.shape[2], tenIn.shape[3]]
)
if self.needs_input_grad[0] == True
else None
)
tenVergrad = (
tenVer.new_empty(
[tenVer.shape[0], tenVer.shape[1], tenVer.shape[2], tenVer.shape[3]]
)
if self.needs_input_grad[1] == True
else None
)
tenHorgrad = (
tenHor.new_empty(
[tenHor.shape[0], tenHor.shape[1], tenHor.shape[2], tenHor.shape[3]]
)
if self.needs_input_grad[2] == True
else None
)
if tenIngrad is not None:
cuda_launch(
cuda_kernel(
"sepconv_ingrad",
sepconv_ingrad,
{
"tenIn": tenIn,
"tenVer": tenVer,
"tenHor": tenHor,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenVergrad": tenVergrad,
"tenHorgrad": tenHorgrad,
},
)
)(
grid=tuple([int((tenIngrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenIngrad.nelement()),
tenIn.data_ptr(),
tenVer.data_ptr(),
tenHor.data_ptr(),
tenOutgrad.data_ptr(),
tenIngrad.data_ptr(),
None,
None,
],
)
# end
if tenVergrad is not None:
cuda_launch(
cuda_kernel(
"sepconv_vergrad",
sepconv_vergrad,
{
"tenIn": tenIn,
"tenVer": tenVer,
"tenHor": tenHor,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenVergrad": tenVergrad,
"tenHorgrad": tenHorgrad,
},
)
)(
grid=tuple([int((tenVergrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenVergrad.nelement()),
tenIn.data_ptr(),
tenVer.data_ptr(),
tenHor.data_ptr(),
tenOutgrad.data_ptr(),
None,
tenVergrad.data_ptr(),
None,
],
)
# end
if tenHorgrad is not None:
cuda_launch(
cuda_kernel(
"sepconv_horgrad",
sepconv_horgrad,
{
"tenIn": tenIn,
"tenVer": tenVer,
"tenHor": tenHor,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenVergrad": tenVergrad,
"tenHorgrad": tenHorgrad,
},
)
)(
grid=tuple([int((tenHorgrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenHorgrad.nelement()),
tenIn.data_ptr(),
tenVer.data_ptr(),
tenHor.data_ptr(),
tenOutgrad.data_ptr(),
None,
None,
tenHorgrad.data_ptr(),
],
)
# end
return tenIngrad, tenVergrad, tenHorgrad
# end
# end
__all__ = ["sepconv_func"]
@@ -1,440 +0,0 @@
import torch
from .utils import cuda_launch, cuda_kernel, cuda_int32
import cupy
import collections
softsplat_flowgrad = """
extern "C" __global__ void __launch_bounds__(512) softsplat_flowgrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenFlowgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenFlowgrad) / SIZE_2(tenFlowgrad) / SIZE_1(tenFlowgrad) ) % SIZE_0(tenFlowgrad);
const int intC = ( intIndex / SIZE_3(tenFlowgrad) / SIZE_2(tenFlowgrad) ) % SIZE_1(tenFlowgrad);
const int intY = ( intIndex / SIZE_3(tenFlowgrad) ) % SIZE_2(tenFlowgrad);
const int intX = ( intIndex ) % SIZE_3(tenFlowgrad);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltFlowgrad = 0.0f;
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = 0.0f;
{{type}} fltNortheast = 0.0f;
{{type}} fltSouthwest = 0.0f;
{{type}} fltSoutheast = 0.0f;
if (intC == 0) {
fltNorthwest = (({{type}}) (-1.0f)) * (({{type}}) (intSoutheastY) - fltY);
fltNortheast = (({{type}}) (+1.0f)) * (({{type}}) (intSouthwestY) - fltY);
fltSouthwest = (({{type}}) (-1.0f)) * (fltY - ({{type}}) (intNortheastY));
fltSoutheast = (({{type}}) (+1.0f)) * (fltY - ({{type}}) (intNorthwestY));
} else if (intC == 1) {
fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (-1.0f));
fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (-1.0f));
fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (({{type}}) (+1.0f));
fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (({{type}}) (+1.0f));
}
for (int intChannel = 0; intChannel < SIZE_1(tenOutgrad); intChannel += 1) {
{{type}} fltIn = VALUE_4(tenIn, intN, intChannel, intY, intX);
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOutgrad)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intNorthwestY, intNorthwestX) * fltIn * fltNorthwest;
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOutgrad)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intNortheastY, intNortheastX) * fltIn * fltNortheast;
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOutgrad)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intSouthwestY, intSouthwestX) * fltIn * fltSouthwest;
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOutgrad)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intSoutheastY, intSoutheastX) * fltIn * fltSoutheast;
}
}
tenFlowgrad[intIndex] = fltFlowgrad;
} }
"""
softsplat_ingrad = """
extern "C" __global__ void __launch_bounds__(512) softsplat_ingrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenFlowgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) / SIZE_1(tenIngrad) ) % SIZE_0(tenIngrad);
const int intC = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) ) % SIZE_1(tenIngrad);
const int intY = ( intIndex / SIZE_3(tenIngrad) ) % SIZE_2(tenIngrad);
const int intX = ( intIndex ) % SIZE_3(tenIngrad);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltIngrad = 0.0f;
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (intSoutheastY) - fltY);
{{type}} fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (intSouthwestY) - fltY);
{{type}} fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (fltY - ({{type}}) (intNortheastY));
{{type}} fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (fltY - ({{type}}) (intNorthwestY));
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOutgrad)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intNorthwestY, intNorthwestX) * fltNorthwest;
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOutgrad)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intNortheastY, intNortheastX) * fltNortheast;
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOutgrad)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intSouthwestY, intSouthwestX) * fltSouthwest;
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOutgrad)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intSoutheastY, intSoutheastX) * fltSoutheast;
}
tenIngrad[intIndex] = fltIngrad;
} }
"""
softsplat_out = """
extern "C" __global__ void __launch_bounds__(512) softsplat_out(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
{{type}}* __restrict__ tenOut
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) / SIZE_1(tenOut) ) % SIZE_0(tenOut);
const int intC = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) ) % SIZE_1(tenOut);
const int intY = ( intIndex / SIZE_3(tenOut) ) % SIZE_2(tenOut);
const int intX = ( intIndex ) % SIZE_3(tenOut);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
{{type}} fltIn = VALUE_4(tenIn, intN, intC, intY, intX);
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (intSoutheastY) - fltY);
{{type}} fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (intSouthwestY) - fltY);
{{type}} fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (fltY - ({{type}}) (intNortheastY));
{{type}} fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (fltY - ({{type}}) (intNorthwestY));
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOut)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intNorthwestY, intNorthwestX)], fltIn * fltNorthwest);
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOut)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intNortheastY, intNortheastX)], fltIn * fltNortheast);
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOut)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intSouthwestY, intSouthwestX)], fltIn * fltSouthwest);
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOut)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intSoutheastY, intSoutheastX)], fltIn * fltSoutheast);
}
} }
"""
# end
class softsplat_func(torch.autograd.Function):
@staticmethod
@torch.cuda.amp.custom_fwd(cast_inputs=torch.float32)
def forward(self, tenIn, tenFlow):
tenOut = tenIn.new_zeros(
[tenIn.shape[0], tenIn.shape[1], tenIn.shape[2], tenIn.shape[3]]
)
if tenIn.is_cuda == True:
cuda_launch(
cuda_kernel(
"softsplat_out",
softsplat_out,
{"tenIn": tenIn, "tenFlow": tenFlow, "tenOut": tenOut},
)
)(
grid=tuple([int((tenOut.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenOut.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOut.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
elif tenIn.is_cuda != True:
assert False
# end
self.save_for_backward(tenIn, tenFlow)
return tenOut
# end
@staticmethod
@torch.cuda.amp.custom_bwd
def backward(self, tenOutgrad):
tenIn, tenFlow = self.saved_tensors
tenOutgrad = tenOutgrad.contiguous()
assert tenOutgrad.is_cuda == True
tenIngrad = (
tenIn.new_zeros(
[tenIn.shape[0], tenIn.shape[1], tenIn.shape[2], tenIn.shape[3]]
)
if self.needs_input_grad[0] == True
else None
)
tenFlowgrad = (
tenFlow.new_zeros(
[tenFlow.shape[0], tenFlow.shape[1], tenFlow.shape[2], tenFlow.shape[3]]
)
if self.needs_input_grad[1] == True
else None
)
if tenIngrad is not None:
cuda_launch(
cuda_kernel(
"softsplat_ingrad",
softsplat_ingrad,
{
"tenIn": tenIn,
"tenFlow": tenFlow,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenFlowgrad": tenFlowgrad,
},
)
)(
grid=tuple([int((tenIngrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenIngrad.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOutgrad.data_ptr(),
tenIngrad.data_ptr(),
None,
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
if tenFlowgrad is not None:
cuda_launch(
cuda_kernel(
"softsplat_flowgrad",
softsplat_flowgrad,
{
"tenIn": tenIn,
"tenFlow": tenFlow,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenFlowgrad": tenFlowgrad,
},
)
)(
grid=tuple([int((tenFlowgrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenFlowgrad.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOutgrad.data_ptr(),
None,
tenFlowgrad.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
return tenIngrad, tenFlowgrad
# end
def FunctionSoftsplat(tenInput, tenFlow, tenMetric, strType):
assert tenMetric is None or tenMetric.shape[1] == 1
assert strType in ["summation", "average", "linear", "softmax"]
if strType == "average":
tenInput = torch.cat(
[
tenInput,
tenInput.new_ones(
tenInput.shape[0], 1, tenInput.shape[2], tenInput.shape[3]
),
],
1,
)
elif strType == "linear":
tenInput = torch.cat([tenInput * tenMetric, tenMetric], 1)
elif strType == "softmax":
tenInput = torch.cat([tenInput * tenMetric.exp(), tenMetric.exp()], 1)
# end
tenOutput = softsplat_func.apply(tenInput, tenFlow)
if strType != "summation":
tenNormalize = tenOutput[:, -1:, :, :]
tenNormalize[tenNormalize == 0.0] = 1.0
tenOutput = tenOutput[:, :-1, :, :] / tenNormalize
# end
return tenOutput
# end
class ModuleSoftsplat(torch.nn.Module):
def __init__(self, strType):
super().__init__()
self.strType = strType
# end
def forward(self, tenInput, tenFlow, tenMetric):
return FunctionSoftsplat(tenInput, tenFlow, tenMetric, self.strType)
# end
# end
def softsplat(
tenIn: torch.Tensor, tenFlow: torch.Tensor, tenMetric: torch.Tensor, strMode: str
):
assert strMode.split("-")[0] in ["sum", "avg", "linear", "soft"]
if strMode == "sum":
assert tenMetric is None
if strMode == "avg":
assert tenMetric is None
if strMode.split("-")[0] == "linear":
assert tenMetric is not None
if strMode.split("-")[0] == "soft":
assert tenMetric is not None
if strMode == "avg":
tenIn = torch.cat(
[
tenIn,
tenIn.new_ones([tenIn.shape[0], 1, tenIn.shape[2], tenIn.shape[3]]),
],
1,
)
elif strMode.split("-")[0] == "linear":
tenIn = torch.cat([tenIn * tenMetric, tenMetric], 1)
elif strMode.split("-")[0] == "soft":
tenIn = torch.cat([tenIn * tenMetric.exp(), tenMetric.exp()], 1)
# end
tenOut = softsplat_func.apply(tenIn, tenFlow)
if strMode.split("-")[0] in ["avg", "linear", "soft"]:
tenNormalize = tenOut[:, -1:, :, :]
if len(strMode.split("-")) == 1:
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "addeps":
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "zeroeps":
tenNormalize[tenNormalize == 0.0] = 1.0
elif strMode.split("-")[1] == "clipeps":
tenNormalize = tenNormalize.clip(0.0000001, None)
# end
tenOut = tenOut[:, :-1, :, :] / tenNormalize
# end
return tenOut
# end
__all__ = ["FunctionSoftsplat", "ModuleSoftsplat", "softsplat", "softsplat_func"]
@@ -1,242 +0,0 @@
import cupy
import os
import re
import torch
import typing
from pathlib import Path
import platform
##########################################################
objCudacache = {}
def cuda_int32(intIn: int):
return cupy.int32(intIn)
# end
def cuda_float32(fltIn: float):
return cupy.float32(fltIn)
# end
def cuda_kernel(strFunction: str, strKernel: str, objVariables: typing.Dict, **replace_kwargs):
if "device" not in objCudacache:
objCudacache["device"] = torch.cuda.get_device_name()
# end
strKey = strFunction
for strVariable in objVariables:
objValue = objVariables[strVariable]
strKey += strVariable
if objValue is None:
continue
elif type(objValue) == int:
strKey += str(objValue)
elif type(objValue) == float:
strKey += str(objValue)
elif type(objValue) == bool:
strKey += str(objValue)
elif type(objValue) == str:
strKey += objValue
elif type(objValue) == torch.Tensor:
strKey += str(objValue.dtype)
strKey += str(objValue.shape)
strKey += str(objValue.stride())
elif True:
print(strVariable, type(objValue))
assert False
# end
# end
strKey += objCudacache["device"]
if strKey not in objCudacache:
for strVariable in objVariables:
objValue = objVariables[strVariable]
if objValue is None:
continue
elif type(objValue) == int:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == float:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == bool:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == str:
strKernel = strKernel.replace("{{" + strVariable + "}}", objValue)
elif type(objValue) == torch.Tensor and objValue.dtype == torch.uint8:
strKernel = strKernel.replace("{{type}}", "unsigned char")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float16:
strKernel = strKernel.replace("{{type}}", "half")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float32:
strKernel = strKernel.replace("{{type}}", "float")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float64:
strKernel = strKernel.replace("{{type}}", "double")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.int32:
strKernel = strKernel.replace("{{type}}", "int")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.int64:
strKernel = strKernel.replace("{{type}}", "long")
elif type(objValue) == torch.Tensor:
print(strVariable, objValue.dtype)
assert False
elif True:
print(strVariable, type(objValue))
assert False
# end
# end
while True:
objMatch = re.search("(SIZE_)([0-4])(\()([^\)]*)(\))", strKernel)
if objMatch is None:
break
# end
intArg = int(objMatch.group(2))
strTensor = objMatch.group(4)
intSizes = objVariables[strTensor].size()
strKernel = strKernel.replace(objMatch.group(), str(intSizes[intArg]))
# end
while True:
objMatch = re.search("(OFFSET_)([0-4])(\()([^\)]+)(\))", strKernel)
if objMatch is None:
break
# end
intArgs = int(objMatch.group(2))
strArgs = objMatch.group(4).split(",")
strTensor = strArgs[0]
intStrides = objVariables[strTensor].stride()
strIndex = [
"(("
+ strArgs[intArg + 1].replace("{", "(").replace("}", ")").strip()
+ ")*"
+ str(intStrides[intArg])
+ ")"
for intArg in range(intArgs)
]
strKernel = strKernel.replace(
objMatch.group(0), "(" + str.join("+", strIndex) + ")"
)
# end
while True:
objMatch = re.search("(VALUE_)([0-4])(\()", strKernel)
if objMatch is None:
break
# end
intStart = objMatch.span()[1]
intStop = objMatch.span()[1]
intParentheses = 1
while True:
intParentheses += 1 if strKernel[intStop] == "(" else 0
intParentheses -= 1 if strKernel[intStop] == ")" else 0
if intParentheses == 0:
break
# end
intStop += 1
# end
intArgs = int(objMatch.group(2))
strArgs = strKernel[intStart:intStop].split(",")
assert intArgs == len(strArgs) - 1
strTensor = strArgs[0]
intStrides = objVariables[strTensor].stride()
strIndex = []
for intArg in range(intArgs):
strIndex.append(
"(("
+ strArgs[intArg + 1].replace("{", "(").replace("}", ")").strip()
+ ")*"
+ str(intStrides[intArg])
+ ")"
)
# end
strKernel = strKernel.replace(
"VALUE_" + str(intArgs) + "(" + strKernel[intStart:intStop] + ")",
strTensor + "[" + str.join("+", strIndex) + "]",
)
# end
for replace_key, value in replace_kwargs.items():
strKernel = strKernel.replace(replace_key, value)
objCudacache[strKey] = {"strFunction": strFunction, "strKernel": strKernel}
# end
return strKey
# end
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
@cupy.memoize(for_each_device=True)
def cuda_launch(strKey: str):
if True:#"CUDA_HOME" not in os.environ:
cuda_home = get_cuda_home_path()
if cuda_home is not None:
os.environ["CUDA_HOME"] = cuda_home
os.environ["CUDA_PATH"] = cuda_home
else:
os.environ["CUDA_HOME"] = "/usr/local/cuda/"
os.environ["CUDA_PATH"] = "/usr/local/cuda/"
# print(objCudacache[strKey]['strKernel'])
# return cupy.cuda.compile_with_cache(objCudacache[strKey]['strKernel'], tuple(['-I ' + os.environ['CUDA_HOME'], '-I ' + os.environ['CUDA_HOME'] + '/include'])).get_function(objCudacache[strKey]['strFunction'])
return cupy.RawModule(code=objCudacache[strKey]["strKernel"]).get_function(
objCudacache[strKey]["strFunction"]
)
@@ -1,150 +0,0 @@
import comfy.model_management as model_management
import torch
import torch.multiprocessing as mp
from .worker_process import f
from .utils import to_shared_memory
parent_conn, child_conn, process = None, None, None
device = model_management.get_torch_device()
def req_to_taichi_process(op_name, *tensors):
global parent_conn, child_conn, process
if parent_conn is None:
mp.set_start_method('spawn', force=True)
parent_conn, child_conn = mp.Pipe()
process = mp.Process(target=f, args=(child_conn, device))
process.start()
tensors = to_shared_memory(tensors)
parent_conn.send((op_name, tensors))
result = parent_conn.recv()
del tensors
if type(result) not in [tuple, list]:
raise Exception(result)
return [tensor.to(device) for tensor in result]
def softsplat(
tenIn: torch.Tensor, tenFlow: torch.Tensor, tenMetric: torch.Tensor, strMode: str
):
assert strMode.split("-")[0] in ["sum", "avg", "linear", "soft"]
if strMode == "sum":
assert tenMetric is None
if strMode == "avg":
assert tenMetric is None
if strMode.split("-")[0] == "linear":
assert tenMetric is not None
if strMode.split("-")[0] == "soft":
assert tenMetric is not None
if strMode == "avg":
tenIn = torch.cat(
[
tenIn,
tenIn.new_ones([tenIn.shape[0], 1, tenIn.shape[2], tenIn.shape[3]]),
],
1,
)
elif strMode.split("-")[0] == "linear":
tenIn = torch.cat([tenIn * tenMetric, tenMetric], 1)
elif strMode.split("-")[0] == "soft":
tenIn = torch.cat([tenIn * tenMetric.exp(), tenMetric.exp()], 1)
# end
tenOut = req_to_taichi_process("softsplat_out", tenIn, tenFlow)[0]
if strMode.split("-")[0] in ["avg", "linear", "soft"]:
tenNormalize = tenOut[:, -1:, :, :]
if len(strMode.split("-")) == 1:
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "addeps":
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "zeroeps":
tenNormalize[tenNormalize == 0.0] = 1.0
elif strMode.split("-")[1] == "clipeps":
tenNormalize = tenNormalize.clip(0.0000001, None)
# end
tenOut = tenOut[:, :-1, :, :] / tenNormalize
# end
return tenOut
def FunctionSoftsplat(tenInput, tenFlow, tenMetric, strType):
assert tenMetric is None or tenMetric.shape[1] == 1
assert strType in ["summation", "average", "linear", "softmax"]
if strType == "average":
tenInput = torch.cat(
[
tenInput,
tenInput.new_ones(
tenInput.shape[0], 1, tenInput.shape[2], tenInput.shape[3]
),
],
1,
)
elif strType == "linear":
tenInput = torch.cat([tenInput * tenMetric, tenMetric], 1)
elif strType == "softmax":
tenInput = torch.cat([tenInput * tenMetric.exp(), tenMetric.exp()], 1)
# end
tenOutput = req_to_taichi_process("softsplat_out", tenInput, tenFlow)[0]
if strType != "summation":
tenNormalize = tenOutput[:, -1:, :, :]
tenNormalize[tenNormalize == 0.0] = 1.0
tenOutput = tenOutput[:, :-1, :, :] / tenNormalize
# end
return tenOutput
# end
class ModuleSoftsplat(torch.nn.Module):
def __init__(self, strType):
super(self).__init__()
self.strType = strType
# end
def forward(self, tenInput, tenFlow, tenMetric):
return FunctionSoftsplat(tenInput, tenFlow, tenMetric, self.strType)
def softsplat_func(tenIn, tenFlow):
return req_to_taichi_process("softsplat_out", tenIn, tenFlow)[0]
class costvol_func:
@staticmethod
def apply(tenOne, tenTwo):
return req_to_taichi_process("costvol_out", tenOne, tenTwo)[0]
class sepconv_func:
@staticmethod
def apply(tenIn, tenVer, tenHor):
return req_to_taichi_process("sepconv_out", tenIn, tenVer, tenHor)[0]
def init():
one_sample = torch.ones(1, 3, 16, 16, dtype=torch.float32, device=device)
softsplat_func(one_sample, one_sample)
costvol_func.apply(one_sample, one_sample)
sepconv_func.apply(one_sample, one_sample, one_sample)
@@ -1,6 +0,0 @@
import torch
class FunctionAdaCoF(torch.autograd.Function):
# end
@staticmethod
def forward(ctx, input, weight, offset_i, offset_j, dilation):
raise NotImplementedError()
@@ -1,2 +0,0 @@
def batch_edt(img, block=1024):
raise NotImplementedError()
@@ -1,15 +0,0 @@
import torch
class _FunctionCorrelation(torch.autograd.Function):
@staticmethod
def forward(self, first, second):
raise NotImplementedError()
def FunctionCorrelation(tenFirst, tenSecond):
raise NotImplementedError()
return _FunctionCorrelation.apply(tenFirst, tenSecond)
class ModuleCorrelation(torch.nn.Module):
def __init__(self):
raise NotImplementedError()
super(ModuleCorrelation, self).__init__()
@@ -1,26 +0,0 @@
import taichi as ti
import taichi.math as tm
""" @ti.kernel
def costvol_out(tenOne: ti.types.ndarray(), tltOne: ti.types.ndarray(), tenTwo: ti.types.ndarray(), tenOut: ti.types.ndarray()):
N, C, H, W = tenOut.shape
for i, ch, y, x in ti.ndrange(N, C, H, W):
for intValue in range(tenOne.shape[1]):
tltOne[intValue] = tenOne[i, intValue, y, x]
tenOut_ch = 0
for intOy in range(y - 4, y + 4 + 1):
for intOx in range(x - 4, x + 4 + 1):
point = tm.ivec2(intOx, intOy)
fltValue = 0.0
for intValue in range(ch):
if (point.y >= 0) and (point.y < H) and (point.x >= 0) and (point.x < W):
fltValue += ti.abs(tltOne[intValue] - tenTwo[i, intValue, point.y, point.x])
else:
fltValue += ti.abs(tltOne[intValue])
tenOut[i, tenOut_ch, y, x] = fltValue / tenOne.shape[1]
tenOut_ch += 1 """
def worker_interface(op_name, tensors):
raise NotImplementedError(op_name)
@@ -1,126 +0,0 @@
#Seperate taichi kernels to another file so that comfy.model_management won't be called in the new process
import taichi as ti
import taichi.math as tm
@ti.func
def put_to_tenOut(tenOut: ti.types.ndarray(), fltIn: ti.i32, flt: ti.i32, pos:tm.uvec2, i:ti.i32, ch:ti.i32):
N, C, H, W = tenOut.shape
if (pos.x >= 0) and (pos.x < W) and (pos.y >= 0) and (pos.y < H):
tenOut[i, ch, pos.y, pos.x] += fltIn * flt
@ti.kernel
def softsplat_out(tenIn: ti.types.ndarray(), tenFlow: ti.types.ndarray(), tenOut: ti.types.ndarray()):
N, C, H, W = tenIn.shape
for i, ch, y, x in ti.ndrange(N, C, H, W):
fltX = x + tenFlow[i, 0, y, x]
fltY = y + tenFlow[i, 1, y, x]
fltIn = tenIn[i, ch, y, x]
northWest = tm.ivec2(ti.floor(fltX), ti.floor(fltY))
northEast = northWest + [1, 0]
southWest = northWest + [0, 1]
southEast = northWest + [1, 1]
fltNorthwest = (southEast.x - fltX) * (southEast.y - fltY)
fltNortheast = (fltX - southWest.x) * (southWest.y - fltY)
fltSouthwest = (northEast.x - fltX) * (fltY - northEast.y)
fltSoutheast = (fltX - northWest.x) * (fltY - northWest.y)
put_to_tenOut(tenOut, fltIn, fltNorthwest, northWest, i, ch)
put_to_tenOut(tenOut, fltIn, fltNortheast, northEast, i, ch)
put_to_tenOut(tenOut, fltIn, fltSouthwest, southWest, i, ch)
put_to_tenOut(tenOut, fltIn, fltSoutheast, southEast, i, ch)
@ti.func
def add_to_fltFlowgrad(fltFlowgrad, tenOutgrad, fltIn, flt, pos, i, ch):
N, C, H, W = tenOutgrad.shape
if (pos.x >= 0) and (pos.x < W) and (pos.y >= 0) and (pos.y < H):
fltFlowgrad += tenOutgrad[i, ch, pos.y, pos.x] * fltIn * flt
@ti.kernel
def softsplat_flowgrad(
tenIn: ti.types.ndarray(),
tenFlow: ti.types.ndarray(),
tenOutgrad: ti.types.ndarray(),
tenIngrad: ti.types.ndarray(),
tenFlowgrad: ti.types.ndarray()
):
N, C, H, W = tenFlowgrad.shape
for i, ch, y, x in ti.ndrange(N, C, H, W):
fltFlowgrad = 0.0
fltX = x + tenFlow[i, 0, y, x]
fltY = y + tenFlow[i, 1, y, x]
northWest = tm.vec2(ti.floor(fltX, dtype=ti.i32), ti.floor(fltY, dtype=ti.i32))
northEast = tm.vec2(northWest.x + 1, northWest.y)
southWest = tm.vec2(northWest.x, northWest.y + 1)
southEast = tm.vec2(northWest.x + 1, northWest.y + 1)
if ch == 0:
fltNorthwest = -1.0 * (southEast.y - fltY)
fltNortheast = +1.0 * (southWest.y - fltY)
fltSouthwest = -1.0 * (fltY - northEast.y)
fltSoutheast = +1.0 * (fltY - northWest.y)
elif ch == 1:
fltNorthwest = -1.0 * (southEast.x - fltX)
fltNortheast = -1.0 * (fltX - southWest.x)
fltSouthwest = +1.0 * (northEast.x - fltX)
fltSoutheast = +1.0 * (fltX - northWest.x)
for outgrad_ch in ti.ndrange(tenOutgrad.shape[1]):
fltIn = tenIn[i, outgrad_ch, y, x]
add_to_fltFlowgrad(fltFlowgrad, tenOutgrad, fltIn, fltNorthwest, northWest, i, outgrad_ch)
add_to_fltFlowgrad(fltFlowgrad, tenOutgrad, fltIn, fltNortheast, northEast, i, outgrad_ch)
add_to_fltFlowgrad(fltFlowgrad, tenOutgrad, fltIn, fltSouthwest, southWest, i, outgrad_ch)
add_to_fltFlowgrad(fltFlowgrad, tenOutgrad, fltIn, fltSoutheast, southEast, i, outgrad_ch)
tenFlowgrad[i] = fltFlowgrad #Is 'i' the same as intIndex?
@ti.func
def add_to_fltIngrad(fltIngrad, tenOutgrad, flt, pos, i, ch):
N, C, H, W = tenOutgrad.shape
if (pos.x >= 0) and (pos.x < W) and (pos.y >= 0) and (pos.y < H):
fltIngrad += tenOutgrad[i, ch, pos.y, pos.x] * flt
@ti.kernel
def softsplat_ingrad(
tenIn: ti.types.ndarray(),
tenFlow: ti.types.ndarray(),
tenOutgrad: ti.types.ndarray(),
tenIngrad: ti.types.ndarray(),
tenFlowgrad: ti.types.ndarray()
):
N, C, H, W = tenIngrad.shape
for i, ch, y, x in ti.ndrange(N, C, H, W):
fltIngrad = 0.0
fltX = x + tenFlow[i, 0, y, x]
fltY = y + tenFlow[i, 1, y, x]
northWest = tm.vec2(ti.floor(fltX, dtype=ti.i32), ti.floor(fltY, dtype=ti.i32))
northEast = tm.vec2(northWest.x + 1, northWest.y)
southWest = tm.vec2(northWest.x, northWest.y + 1)
southEast = tm.vec2(northWest.x + 1, northWest.y + 1)
fltNorthwest = (southEast.x - fltX) * (southEast.y - fltY)
fltNortheast = (fltX - southWest.x) * (southWest.y - fltY)
fltSouthwest = (northEast.x - fltX) * (fltY - northEast.y)
fltSoutheast = (fltX - northWest.x) * (fltY - northWest.y)
add_to_fltIngrad(fltIngrad, tenOutgrad, fltNorthwest, northWest, i, ch)
add_to_fltIngrad(fltIngrad, tenOutgrad, fltNortheast, northEast, i, ch)
add_to_fltIngrad(fltIngrad, tenOutgrad, fltSouthwest, southWest, i, ch)
add_to_fltIngrad(fltIngrad, tenOutgrad, fltSoutheast, southEast, i, ch)
tenIngrad[i] = fltIngrad
# end
def worker_interface(op_name, tensors):
if op_name == "softsplat_out":
tenIn, tenFlow = tensors
tenOut = tenIn.new_zeros(tenIn.shape)
softsplat_out(tenIn, tenFlow, tenOut)
return (tenOut, )
raise NotImplementedError(op_name)
__all__ = ["worker_interface"]
@@ -1,39 +0,0 @@
import taichi as ti
import taichi.math as tm
from functools import reduce
@ti.kernel
def sepconv_out(tenIn: ti.types.ndarray(), tenVer: ti.types.ndarray(), tenHor: ti.types.ndarray(), tenOut: ti.types.ndarray()):
N, C, H, W = tenIn.shape
intIndex = 0
for i, ch, y, x in ti.ndrange(N, C, H, W):
fltOut, fltKahanc, fltKahany, fltKahant = 0.0, 0.0, 0.0, 0.0
for intFy, intFx in ti.ndrange(tenVer.shape[1], tenHor.shape[1]):
fltKahany = tenIn[i, ch, y + intFy, x + intFx] * tenVer[i, intFy, y, x] * tenHor[i, intFx, y, x]
fltKahany = fltKahany - fltKahanc
fltKahant = fltOut + fltKahany
fltKahanc = (fltKahant - fltOut) - fltKahany
fltOut = fltKahant
tenOut[intIndex] = fltOut
intIndex += 1
def worker_interface(op_name, tensors):
if op_name == "sepconv_out":
tenIn, tenVer, tenHor = tensors
real_tenOut_shape = [
tenIn.shape[0],
tenIn.shape[1],
tenVer.shape[2] and tenHor.shape[2],
tenVer.shape[3] and tenHor.shape[3],
]
tenOut = tenIn.new_zeros([
int(reduce(lambda a, b: a * b, real_tenOut_shape))
])
sepconv_out(tenIn, tenVer, tenHor, tenOut)
tenOut = tenOut.view(*real_tenOut_shape)
return (tenOut, )
raise NotImplementedError(op_name)
__all__ = ["worker_interface"]
@@ -1,11 +0,0 @@
import platform
import torch
def to_shared_memory(tensors: tuple[torch.Tensor]):
return [tensor.cpu() for tensor in tensors if tensor is not None]
""" if platform.system() == "Windows":
return [tensor.cpu() for tensor in tensors if tensor is not None]
return [tensor.share_memory_() for tensor in tensors if tensor is not None] """
def to_device(tensors: tuple[torch.Tensor], device: torch.device):
return [tensor.to(device) for tensor in tensors if tensor is not None]
@@ -1,26 +0,0 @@
import torch.multiprocessing as mp
import torch
from .raw_softsplat import worker_interface as raw_softsplat
from .costvol import worker_interface as costvol
from .sepconv import worker_interface as sepconv
from .utils import to_shared_memory, to_device
import taichi as ti
import traceback
def f(child_conn, device: torch.DeviceObjType):
ti.init(arch=ti.gpu)
while True:
op_name, tensors = child_conn.recv()
tensors = to_device(tensors, device)
try:
if "softsplat" in op_name:
result = raw_softsplat(op_name, tensors)
elif "costvol" in op_name:
result = costvol(op_name, tensors)
elif "sepconv" in op_name:
result = sepconv(op_name, tensors)
else:
raise NotImplementedError(op_name)
child_conn.send(to_shared_memory(result))
except:
child_conn.send(traceback.format_exc())
@@ -1,107 +0,0 @@
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, generic_frame_loop, InterpolationStateList
import typing
from comfy.model_management import get_torch_device
import re
from functools import cmp_to_key
from packaging import version
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAME_VER_DICT = {
"rife40.pth": "4.0",
"rife41.pth": "4.0",
"rife42.pth": "4.2",
"rife43.pth": "4.3",
"rife44.pth": "4.3",
"rife45.pth": "4.5",
"rife46.pth": "4.6",
"rife47.pth": "4.7",
"rife48.pth": "4.7",
"rife49.pth": "4.7",
"sudo_rife4_269.662_testV1_scale1.pth": "4.0"
#Arch 4.10 doesn't work due to state dict mismatch
#TODO: Investigating and fix it
#"rife410.pth": "4.10",
#"rife411.pth": "4.10",
#"rife412.pth": "4.10"
}
class RIFE_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (
sorted(list(CKPT_NAME_VER_DICT.keys()), key=lambda ckpt_name: version.parse(CKPT_NAME_VER_DICT[ckpt_name])),
{"default": "rife47.pth"}
),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 1}),
"fast_mode": ("BOOLEAN", {"default":True}),
"ensemble": ("BOOLEAN", {"default":True}),
"scale_factor": ([0.25, 0.5, 1.0, 2.0, 4.0], {"default": 1.0})
},
"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,
fast_mode = False,
ensemble = False,
scale_factor = 1.0,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
"""
Perform video frame interpolation using a given checkpoint model.
Args:
ckpt_name (str): The name of the checkpoint model to use.
frames (torch.Tensor): A tensor containing input video frames.
clear_cache_after_n_frames (int, optional): The number of frames to process before clearing CUDA cache
to prevent memory overflow. Defaults to 10. Lower numbers are safer but mean more processing time.
How high you should set it depends on how many input frames there are, input resolution (after upscaling),
how many times you want to multiply them, and how long you're willing to wait for the process to complete.
multiplier (int, optional): The multiplier for each input frame. 60 input frames * 2 = 120 output frames. Defaults to 2.
Returns:
tuple: A tuple containing the output interpolated frames.
Note:
This method interpolates frames in a video sequence using a specified checkpoint model.
It processes each frame sequentially, generating interpolated frames between them.
To prevent memory overflow, it clears the CUDA cache after processing a specified number of frames.
"""
from .rife_arch import IFNet
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
arch_ver = CKPT_NAME_VER_DICT[ckpt_name]
interpolation_model = IFNet(arch_ver=arch_ver)
interpolation_model.load_state_dict(torch.load(model_path))
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model, scale_list, in_fast_mode, in_ensemble):
return model(frame_0, frame_1, timestep, scale_list, in_fast_mode, in_ensemble)
scale_list = [8 / scale_factor, 4 / scale_factor, 2 / scale_factor, 1 / scale_factor]
args = [interpolation_model, scale_list, fast_mode, ensemble]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, dtype=torch.float32)
)
return (out,)
@@ -1,581 +0,0 @@
"""
26-Dez-21
https://github.com/hzwer/Practical-RIFE
https://github.com/hzwer/Practical-RIFE/blob/main/model/warplayer.py
https://github.com/HolyWu/vs-rife/blob/master/vsrife/__init__.py
"""
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.optim import AdamW
import torch
import torch.nn.functional as F
import torch.nn as nn
import torch.optim as optim
import warnings
from comfy.model_management import get_torch_device
device = get_torch_device()
backwarp_tenGrid = {}
class ResConv(nn.Module):
def __init__(self, c, dilation=1):
super(ResConv, self).__init__()
self.conv = nn.Conv2d(c, c, 3, 1, dilation, dilation=dilation, groups=1)
self.beta = nn.Parameter(torch.ones((1, c, 1, 1)), requires_grad=True)
self.relu = nn.LeakyReLU(0.2, True)
def forward(self, x):
return self.relu(self.conv(x) * self.beta + x)
def warp(tenInput, tenFlow):
k = (str(tenFlow.device), str(tenFlow.size()))
if k not in backwarp_tenGrid:
tenHorizontal = (
torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device)
.view(1, 1, 1, tenFlow.shape[3])
.expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)
)
tenVertical = (
torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device)
.view(1, 1, tenFlow.shape[2], 1)
.expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])
)
backwarp_tenGrid[k] = torch.cat([tenHorizontal, tenVertical], 1).to(device)
tenFlow = torch.cat(
[
tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),
tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0),
],
1,
)
g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)
if tenInput.type() == "torch.cuda.HalfTensor":
g = g.half()
return torch.nn.functional.grid_sample(
input=tenInput,
grid=g,
mode="bilinear",
padding_mode="border",
align_corners=True,
)
def conv(
in_planes,
out_planes,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
arch_ver="4.0",
):
if arch_ver == "4.0":
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
),
nn.PReLU(out_planes),
)
if arch_ver in ["4.2", "4.3", "4.5", "4.6", "4.7", "4.10"]:
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
),
nn.LeakyReLU(0.2, True),
)
def conv_woact(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
),
)
def conv_woact(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(
in_planes,
out_planes,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=True,
)
)
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1, arch_ver="4.0"):
if arch_ver == "4.0":
return nn.Sequential(
torch.nn.ConvTranspose2d(
in_channels=in_planes,
out_channels=out_planes,
kernel_size=4,
stride=2,
padding=1,
bias=True,
),
nn.PReLU(out_planes),
)
if arch_ver in ["4.2", "4.3", "4.5", "4.6", "4.7", "4.10"]:
return nn.Sequential(
torch.nn.ConvTranspose2d(
in_channels=in_planes,
out_channels=out_planes,
kernel_size=4,
stride=2,
padding=1,
bias=True,
),
nn.LeakyReLU(0.2, True),
)
class Conv2(nn.Module):
def __init__(self, in_planes, out_planes, stride=2, arch_ver="4.0"):
super(Conv2, self).__init__()
self.conv1 = conv(in_planes, out_planes, 3, stride, 1, arch_ver=arch_ver)
self.conv2 = conv(out_planes, out_planes, 3, 1, 1, arch_ver=arch_ver)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
return x
class IFBlock(nn.Module):
def __init__(self, in_planes, c=64, arch_ver="4.0"):
super(IFBlock, self).__init__()
self.arch_ver = arch_ver
self.conv0 = nn.Sequential(
conv(in_planes, c // 2, 3, 2, 1, arch_ver=arch_ver),
conv(c // 2, c, 3, 2, 1, arch_ver=arch_ver),
)
self.arch_ver = arch_ver
if arch_ver in ["4.0", "4.2", "4.3"]:
self.convblock = nn.Sequential(
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
conv(c, c, arch_ver=arch_ver),
)
self.lastconv = nn.ConvTranspose2d(c, 5, 4, 2, 1)
if arch_ver in ["4.5", "4.6", "4.7", "4.10"]:
self.convblock = nn.Sequential(
ResConv(c),
ResConv(c),
ResConv(c),
ResConv(c),
ResConv(c),
ResConv(c),
ResConv(c),
ResConv(c),
)
if arch_ver == "4.5":
self.lastconv = nn.Sequential(
nn.ConvTranspose2d(c, 4 * 5, 4, 2, 1), nn.PixelShuffle(2)
)
if arch_ver in ["4.6", "4.7", "4.10"]:
self.lastconv = nn.Sequential(
nn.ConvTranspose2d(c, 4 * 6, 4, 2, 1), nn.PixelShuffle(2)
)
def forward(self, x, flow=None, scale=1):
x = F.interpolate(
x, scale_factor=1.0 / scale, mode="bilinear", align_corners=False
)
if flow is not None:
flow = (
F.interpolate(
flow, scale_factor=1.0 / scale, mode="bilinear", align_corners=False
)
* 1.0
/ scale
)
x = torch.cat((x, flow), 1)
feat = self.conv0(x)
if self.arch_ver == "4.0":
feat = self.convblock(feat) + feat
if self.arch_ver in ["4.2", "4.3", "4.5", "4.6", "4.7", "4.10"]:
feat = self.convblock(feat)
tmp = self.lastconv(feat)
if self.arch_ver in ["4.0", "4.2", "4.3"]:
tmp = F.interpolate(
tmp, scale_factor=scale * 2, mode="bilinear", align_corners=False
)
flow = tmp[:, :4] * scale * 2
if self.arch_ver in ["4.5", "4.6", "4.7", "4.10"]:
tmp = F.interpolate(
tmp, scale_factor=scale, mode="bilinear", align_corners=False
)
flow = tmp[:, :4] * scale
mask = tmp[:, 4:5]
return flow, mask
class Contextnet(nn.Module):
def __init__(self, arch_ver="4.0"):
super(Contextnet, self).__init__()
c = 16
self.conv1 = Conv2(3, c, arch_ver=arch_ver)
self.conv2 = Conv2(c, 2 * c, arch_ver=arch_ver)
self.conv3 = Conv2(2 * c, 4 * c, arch_ver=arch_ver)
self.conv4 = Conv2(4 * c, 8 * c, arch_ver=arch_ver)
def forward(self, x, flow):
x = self.conv1(x)
flow = (
F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False)
* 0.5
)
f1 = warp(x, flow)
x = self.conv2(x)
flow = (
F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False)
* 0.5
)
f2 = warp(x, flow)
x = self.conv3(x)
flow = (
F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False)
* 0.5
)
f3 = warp(x, flow)
x = self.conv4(x)
flow = (
F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False)
* 0.5
)
f4 = warp(x, flow)
return [f1, f2, f3, f4]
class Unet(nn.Module):
def __init__(self, arch_ver="4.0"):
super(Unet, self).__init__()
c = 16
self.down0 = Conv2(17, 2 * c, arch_ver=arch_ver)
self.down1 = Conv2(4 * c, 4 * c, arch_ver=arch_ver)
self.down2 = Conv2(8 * c, 8 * c, arch_ver=arch_ver)
self.down3 = Conv2(16 * c, 16 * c, arch_ver=arch_ver)
self.up0 = deconv(32 * c, 8 * c, arch_ver=arch_ver)
self.up1 = deconv(16 * c, 4 * c, arch_ver=arch_ver)
self.up2 = deconv(8 * c, 2 * c, arch_ver=arch_ver)
self.up3 = deconv(4 * c, c, arch_ver=arch_ver)
self.conv = nn.Conv2d(c, 3, 3, 1, 1)
def forward(self, img0, img1, warped_img0, warped_img1, mask, flow, c0, c1):
s0 = self.down0(
torch.cat((img0, img1, warped_img0, warped_img1, mask, flow), 1)
)
s1 = self.down1(torch.cat((s0, c0[0], c1[0]), 1))
s2 = self.down2(torch.cat((s1, c0[1], c1[1]), 1))
s3 = self.down3(torch.cat((s2, c0[2], c1[2]), 1))
x = self.up0(torch.cat((s3, c0[3], c1[3]), 1))
x = self.up1(torch.cat((x, s2), 1))
x = self.up2(torch.cat((x, s1), 1))
x = self.up3(torch.cat((x, s0), 1))
x = self.conv(x)
return torch.sigmoid(x)
"""
currently supports 4.0-4.12
4.0: 4.0, 4.1
4.2: 4.2
4.3: 4.3, 4.4
4.5: 4.5
4.6: 4.6
4.7: 4.7, 4.8, 4.9
4.10: 4.10 4.11 4.12
"""
class IFNet(nn.Module):
def __init__(self, arch_ver="4.0"):
super(IFNet, self).__init__()
self.arch_ver = arch_ver
if arch_ver in ["4.0", "4.2", "4.3", "4.5", "4.6"]:
self.block0 = IFBlock(7, c=192, arch_ver=arch_ver)
self.block1 = IFBlock(8 + 4, c=128, arch_ver=arch_ver)
self.block2 = IFBlock(8 + 4, c=96, arch_ver=arch_ver)
self.block3 = IFBlock(8 + 4, c=64, arch_ver=arch_ver)
if arch_ver in ["4.7"]:
self.block0 = IFBlock(7 + 8, c=192, arch_ver=arch_ver)
self.block1 = IFBlock(8 + 4 + 8, c=128, arch_ver=arch_ver)
self.block2 = IFBlock(8 + 4 + 8, c=96, arch_ver=arch_ver)
self.block3 = IFBlock(8 + 4 + 8, c=64, arch_ver=arch_ver)
self.encode = nn.Sequential(
nn.Conv2d(3, 16, 3, 2, 1), nn.ConvTranspose2d(16, 4, 4, 2, 1)
)
if arch_ver in ["4.10"]:
self.block0 = IFBlock(7 + 16, c=192)
self.block1 = IFBlock(8 + 4 + 16, c=128)
self.block2 = IFBlock(8 + 4 + 16, c=96)
self.block3 = IFBlock(8 + 4 + 16, c=64)
self.encode = nn.Sequential(
nn.Conv2d(3, 32, 3, 2, 1),
nn.LeakyReLU(0.2, True),
nn.Conv2d(32, 32, 3, 1, 1),
nn.LeakyReLU(0.2, True),
nn.Conv2d(32, 32, 3, 1, 1),
nn.LeakyReLU(0.2, True),
nn.ConvTranspose2d(32, 8, 4, 2, 1),
)
if arch_ver in ["4.0", "4.2", "4.3"]:
self.contextnet = Contextnet(arch_ver=arch_ver)
self.unet = Unet(arch_ver=arch_ver)
self.arch_ver = arch_ver
def forward(
self,
img0,
img1,
timestep=0.5,
scale_list=[8, 4, 2, 1],
training=True,
fastmode=True,
ensemble=False,
return_flow=False,
):
img0 = torch.clamp(img0, 0, 1)
img1 = torch.clamp(img1, 0, 1)
n, c, h, w = img0.shape
ph = ((h - 1) // 64 + 1) * 64
pw = ((w - 1) // 64 + 1) * 64
padding = (0, pw - w, 0, ph - h)
img0 = F.pad(img0, padding)
img1 = F.pad(img1, padding)
x = torch.cat((img0, img1), 1)
if training == False:
channel = x.shape[1] // 2
img0 = x[:, :channel]
img1 = x[:, channel:]
if not torch.is_tensor(timestep):
timestep = (x[:, :1].clone() * 0 + 1) * timestep
else:
timestep = timestep.repeat(1, 1, img0.shape[2], img0.shape[3])
flow_list = []
merged = []
mask_list = []
if self.arch_ver in ["4.7", "4.10"]:
f0 = self.encode(img0[:, :3])
f1 = self.encode(img1[:, :3])
warped_img0 = img0
warped_img1 = img1
flow = None
mask = None
block = [self.block0, self.block1, self.block2, self.block3]
for i in range(4):
if flow is None:
# 4.0-4.6
if self.arch_ver in ["4.0", "4.2", "4.3", "4.5", "4.6"]:
flow, mask = block[i](
torch.cat((img0[:, :3], img1[:, :3], timestep), 1),
None,
scale=scale_list[i],
)
if ensemble:
f1, m1 = block[i](
torch.cat((img1[:, :3], img0[:, :3], 1 - timestep), 1),
None,
scale=scale_list[i],
)
flow = (flow + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 2
mask = (mask + (-m1)) / 2
# 4.7+
if self.arch_ver in ["4.7", "4.10"]:
flow, mask = block[i](
torch.cat((img0[:, :3], img1[:, :3], f0, f1, timestep), 1),
None,
scale=scale_list[i],
)
if ensemble:
f_, m_ = block[i](
torch.cat(
(img1[:, :3], img0[:, :3], f1, f0, 1 - timestep), 1
),
None,
scale=scale_list[i],
)
flow = (flow + torch.cat((f_[:, 2:4], f_[:, :2]), 1)) / 2
mask = (mask + (-m_)) / 2
else:
# 4.0-4.6
if self.arch_ver in ["4.0", "4.2", "4.3", "4.5", "4.6"]:
f0, m0 = block[i](
torch.cat(
(warped_img0[:, :3], warped_img1[:, :3], timestep, mask), 1
),
flow,
scale=scale_list[i],
)
if self.arch_ver in ["4.0"]:
if (
i == 1
and f0[:, :2].abs().max() > 32
and f0[:, 2:4].abs().max() > 32
and not training
):
for k in range(4):
scale_list[k] *= 2
flow, mask = block[0](
torch.cat((img0[:, :3], img1[:, :3], timestep), 1),
None,
scale=scale_list[0],
)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
f0, m0 = block[i](
torch.cat(
(
warped_img0[:, :3],
warped_img1[:, :3],
timestep,
mask,
),
1,
),
flow,
scale=scale_list[i],
)
# 4.7+
if self.arch_ver in ["4.7", "4.10"]:
fd, m0 = block[i](
torch.cat(
(
warped_img0[:, :3],
warped_img1[:, :3],
warp(f0, flow[:, :2]),
warp(f1, flow[:, 2:4]),
timestep,
mask,
),
1,
),
flow,
scale=scale_list[i],
)
flow = flow + fd
# 4.0-4.6 ensemble
if ensemble and self.arch_ver in [
"4.0",
"4.2",
"4.3",
"4.5",
"4.6",
]:
f1, m1 = block[i](
torch.cat(
(
warped_img1[:, :3],
warped_img0[:, :3],
1 - timestep,
-mask,
),
1,
),
torch.cat((flow[:, 2:4], flow[:, :2]), 1),
scale=scale_list[i],
)
f0 = (f0 + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 2
m0 = (m0 + (-m1)) / 2
# 4.7+ ensemble
if ensemble and self.arch_ver in ["4.7", "4.10"]:
wf0 = warp(f0, flow[:, :2])
wf1 = warp(f1, flow[:, 2:4])
f_, m_ = block[i](
torch.cat(
(
warped_img1[:, :3],
warped_img0[:, :3],
wf1,
wf0,
1 - timestep,
-mask,
),
1,
),
torch.cat((flow[:, 2:4], flow[:, :2]), 1),
scale=scale_list[i],
)
fd = (fd + torch.cat((f_[:, 2:4], f_[:, :2]), 1)) / 2
mask = (m0 + (-m_)) / 2
if self.arch_ver in ["4.0", "4.2", "4.3", "4.5", "4.6"]:
flow = flow + f0
mask = mask + m0
if not ensemble and self.arch_ver in ["4.7", "4.10"]:
mask = m0
mask_list.append(mask)
flow_list.append(flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
merged.append((warped_img0, warped_img1))
if self.arch_ver in ["4.0", "4.1", "4.2", "4.3", "4.4", "4.5", "4.6"]:
mask_list[3] = torch.sigmoid(mask_list[3])
merged[3] = merged[3][0] * mask_list[3] + merged[3][1] * (1 - mask_list[3])
if self.arch_ver in ["4.7", "4.10"]:
mask = torch.sigmoid(mask)
merged[3] = warped_img0 * mask + warped_img1 * (1 - mask)
if not fastmode and self.arch_ver in ["4.0", "4.2", "4.3"]:
c0 = self.contextnet(img0, flow[:, :2])
c1 = self.contextnet(img1, flow[:, 2:4])
tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)
res = tmp[:, :3] * 2 - 1
merged[3] = torch.clamp(merged[3] + res, 0, 1)
return merged[3][:, :, :h, :w]
@@ -1,56 +0,0 @@
import torch
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames
import typing
from comfy.model_management import soft_empty_cache, get_torch_device
from vfi_utils import InterpolationStateList, generic_frame_loop
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAMES = ["sepconv.pth"]
class SepconvVFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (CKPT_NAMES, ),
"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: InterpolationStateList = None,
**kwargs
):
from .sepconv_enhanced import Network
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
interpolation_model = Network()
interpolation_model.load_state_dict(torch.load(model_path))
interpolation_model.eval().to(get_torch_device())
frames = preprocess_frames(frames)
def return_middle_frame(frame_0, frame_1, timestep, model):
return model(frame_0, frame_1)
args = [interpolation_model]
out = postprocess_frames(
generic_frame_loop(type(self).__name__, frames, clear_cache_after_n_frames, multiplier, return_middle_frame, *args,
interpolation_states=optional_interpolation_states, use_timestep=False, dtype=torch.float32)
)
return (out,)
@@ -1,748 +0,0 @@
"""
23-nov-21
https://github.com/sniklaus/revisiting-sepconv/blob/fea509d98157170df1fb35bf615bd41d98858e1a/run.py
https://github.com/sniklaus/revisiting-sepconv/blob/fea509d98157170df1fb35bf615bd41d98858e1a/sepconv/sepconv.py
Deleted stuffs about arguments_strModel and getopt
"""
#!/usr/bin/env python
import torch
import typing
from comfy.model_management import get_torch_device
##########################################################
from vfi_models.ops import sepconv_func
##########################################################
import torch
import math
import numpy
import os
import PIL
import PIL.Image
import sys
import typing
##########################################################
assert (
int(str("").join(torch.__version__.split(".")[0:2])) >= 13
) # requires at least pytorch version 1.3.0
torch.set_grad_enabled(
False
) # make sure to not compute gradients for computational performance
torch.backends.cudnn.enabled = (
True # make sure to use cudnn for computational performance
)
##########################################################
##########################################################
class Basic(torch.nn.Module):
def __init__(
self,
strType: str,
intChans: typing.List[int],
objScratch: typing.Optional[typing.Dict] = None,
):
super().__init__()
self.strType = strType
self.netEvenize = None
self.netMain = None
self.netShortcut = None
intIn = intChans[0]
intOut = intChans[-1]
netMain = []
intChans = intChans.copy()
fltStride = 1.0
for intPart, strPart in enumerate(self.strType.split("+")[0].split("-")):
if strPart.startswith("conv") == True:
intKsize = 3
intPad = 1
strPad = "zeros"
if "(" in strPart:
intKsize = int(strPart.split("(")[1].split(")")[0].split(",")[0])
intPad = int(math.floor(0.5 * (intKsize - 1)))
if "replpad" in strPart.split("(")[1].split(")")[0].split(","):
strPad = "replicate"
if "reflpad" in strPart.split("(")[1].split(")")[0].split(","):
strPad = "reflect"
# end
if "nopad" in self.strType.split("+"):
intPad = 0
# end
netMain += [
torch.nn.Conv2d(
in_channels=intChans[0],
out_channels=intChans[1],
kernel_size=intKsize,
stride=1,
padding=intPad,
padding_mode=strPad,
bias="nobias" not in self.strType.split("+"),
)
]
intChans = intChans[1:]
fltStride *= 1.0
elif strPart.startswith("sconv") == True:
intKsize = 3
intPad = 1
strPad = "zeros"
if "(" in strPart:
intKsize = int(strPart.split("(")[1].split(")")[0].split(",")[0])
intPad = int(math.floor(0.5 * (intKsize - 1)))
if "replpad" in strPart.split("(")[1].split(")")[0].split(","):
strPad = "replicate"
if "reflpad" in strPart.split("(")[1].split(")")[0].split(","):
strPad = "reflect"
# end
if "nopad" in self.strType.split("+"):
intPad = 0
# end
netMain += [
torch.nn.Conv2d(
in_channels=intChans[0],
out_channels=intChans[1],
kernel_size=intKsize,
stride=2,
padding=intPad,
padding_mode=strPad,
bias="nobias" not in self.strType.split("+"),
)
]
intChans = intChans[1:]
fltStride *= 2.0
elif strPart.startswith("up") == True:
class Up(torch.nn.Module):
def __init__(self, strType):
super().__init__()
self.strType = strType
# end
def forward(self, tenIn: torch.Tensor) -> torch.Tensor:
if self.strType == "nearest":
return torch.nn.functional.interpolate(
input=tenIn,
scale_factor=2.0,
mode="nearest",
align_corners=False,
)
elif self.strType == "bilinear":
return torch.nn.functional.interpolate(
input=tenIn,
scale_factor=2.0,
mode="bilinear",
align_corners=False,
)
elif self.strType == "pyramid":
return pyramid(tenIn, None, "up")
elif self.strType == "shuffle":
return torch.nn.functional.pixel_shuffle(
tenIn, upscale_factor=2
) # https://github.com/pytorch/pytorch/issues/62854
# end
assert False # to make torchscript happy
# end
# end
strType = "bilinear"
if "(" in strPart:
if "nearest" in strPart.split("(")[1].split(")")[0].split(","):
strType = "nearest"
if "pyramid" in strPart.split("(")[1].split(")")[0].split(","):
strType = "pyramid"
if "shuffle" in strPart.split("(")[1].split(")")[0].split(","):
strType = "shuffle"
# end
netMain += [Up(strType)]
fltStride *= 0.5
elif strPart.startswith("prelu") == True:
netMain += [
torch.nn.PReLU(
num_parameters=1,
init=float(strPart.split("(")[1].split(")")[0].split(",")[0]),
)
]
fltStride *= 1.0
elif True:
assert False
# end
# end
self.netMain = torch.nn.Sequential(*netMain)
for strPart in self.strType.split("+")[1:]:
if strPart.startswith("skip") == True:
if intIn == intOut and fltStride == 1.0:
self.netShortcut = torch.nn.Identity()
elif intIn != intOut and fltStride == 1.0:
self.netShortcut = torch.nn.Conv2d(
in_channels=intIn,
out_channels=intOut,
kernel_size=1,
stride=1,
padding=0,
bias="nobias" not in self.strType.split("+"),
)
elif intIn == intOut and fltStride != 1.0:
class Down(torch.nn.Module):
def __init__(self, fltScale):
super().__init__()
self.fltScale = fltScale
# end
def forward(self, tenIn: torch.Tensor) -> torch.Tensor:
return torch.nn.functional.interpolate(
input=tenIn,
scale_factor=self.fltScale,
mode="bilinear",
align_corners=False,
)
# end
# end
self.netShortcut = Down(1.0 / fltStride)
elif intIn != intOut and fltStride != 1.0:
class Down(torch.nn.Module):
def __init__(self, fltScale):
super().__init__()
self.fltScale = fltScale
# end
def forward(self, tenIn: torch.Tensor) -> torch.Tensor:
return torch.nn.functional.interpolate(
input=tenIn,
scale_factor=self.fltScale,
mode="bilinear",
align_corners=False,
)
# end
# end
self.netShortcut = torch.nn.Sequential(
Down(1.0 / fltStride),
torch.nn.Conv2d(
in_channels=intIn,
out_channels=intOut,
kernel_size=1,
stride=1,
padding=0,
bias="nobias" not in self.strType.split("+"),
),
)
# end
elif strPart.startswith("...") == True:
pass
# end
# end
assert len(intChans) == 1
# end
def forward(self, tenIn: torch.Tensor) -> torch.Tensor:
if self.netEvenize is not None:
tenIn = self.netEvenize(tenIn)
# end
tenOut = self.netMain(tenIn)
if self.netShortcut is not None:
tenOut = tenOut + self.netShortcut(tenIn)
# end
return tenOut
# end
# end
class Encode(torch.nn.Module):
objScratch: typing.Dict[str, typing.List[int]] = None
def __init__(
self,
intIns: typing.List[int],
intOuts: typing.List[int],
strHor: str,
strVer: str,
objScratch: typing.Dict[str, typing.List[int]],
):
super().__init__()
assert len(intIns) == len(intOuts)
assert len(intOuts) == len(intIns)
self.intRows = len(intIns) and len(intOuts)
self.intIns = intIns.copy()
self.intOuts = intOuts.copy()
self.strHor = strHor
self.strVer = strVer
self.objScratch = objScratch
self.netHor = torch.nn.ModuleList()
self.netVer = torch.nn.ModuleList()
for intRow in range(self.intRows):
netHor = torch.nn.Identity()
netVer = torch.nn.Identity()
if self.intOuts[intRow] != 0:
if self.intIns[intRow] != 0:
netHor = Basic(
self.strHor,
[
self.intIns[intRow],
self.intOuts[intRow],
self.intOuts[intRow],
],
objScratch,
)
# end
if intRow != 0:
netVer = Basic(
self.strVer,
[
self.intOuts[intRow - 1],
self.intOuts[intRow],
self.intOuts[intRow],
],
objScratch,
)
# end
# end
self.netHor.append(netHor)
self.netVer.append(netVer)
# end
# end
def forward(self, tenIns: typing.List[torch.Tensor]) -> typing.List[torch.Tensor]:
intRow = 0
for netHor in self.netHor:
if self.intOuts[intRow] != 0:
if self.intIns[intRow] != 0:
tenIns[intRow] = netHor(tenIns[intRow])
# end
# end
intRow += 1
# end
intRow = 0
for netVer in self.netVer:
if self.intOuts[intRow] != 0:
if intRow != 0:
tenIns[intRow] = tenIns[intRow] + netVer(tenIns[intRow - 1])
# end
# end
intRow += 1
# end
for intRow, tenIn in enumerate(tenIns):
self.objScratch["levelshape" + str(intRow)] = tenIn.shape
# end
return tenIns
# end
# end
class Decode(torch.nn.Module):
objScratch: typing.Dict[str, typing.List[int]] = None
def __init__(
self,
intIns: typing.List[int],
intOuts: typing.List[int],
strHor: str,
strVer: str,
objScratch: typing.Dict[str, typing.List[int]],
):
super().__init__()
assert len(intIns) == len(intOuts)
assert len(intOuts) == len(intIns)
self.intRows = len(intIns) and len(intOuts)
self.intIns = intIns.copy()
self.intOuts = intOuts.copy()
self.strHor = strHor
self.strVer = strVer
self.objScratch = objScratch
self.netHor = torch.nn.ModuleList()
self.netVer = torch.nn.ModuleList()
for intRow in range(self.intRows - 1, -1, -1):
netHor = torch.nn.Identity()
netVer = torch.nn.Identity()
if self.intOuts[intRow] != 0:
if self.intIns[intRow] != 0:
netHor = Basic(
self.strHor,
[
self.intIns[intRow],
self.intOuts[intRow],
self.intOuts[intRow],
],
objScratch,
)
# end
if intRow != self.intRows - 1:
netVer = Basic(
self.strVer,
[
self.intOuts[intRow + 1],
self.intOuts[intRow],
self.intOuts[intRow],
],
objScratch,
)
# end
# end
self.netHor.append(netHor)
self.netVer.append(netVer)
# end
# end
def forward(self, tenIns: typing.List[torch.Tensor]) -> typing.List[torch.Tensor]:
intRow = self.intRows - 1
for netHor in self.netHor:
if self.intOuts[intRow] != 0:
if self.intIns[intRow] != 0:
tenIns[intRow] = netHor(tenIns[intRow])
# end
# end
intRow -= 1
# end
intRow = self.intRows - 1
for netVer in self.netVer:
if self.intOuts[intRow] != 0:
if intRow != self.intRows - 1:
tenVer = netVer(tenIns[intRow + 1])
if "levelshape" + str(intRow) in self.objScratch:
if (
tenVer.shape[2]
== self.objScratch["levelshape" + str(intRow)][2] + 1
):
tenVer = torch.nn.functional.pad(
input=tenVer,
pad=[0, 0, 0, -1],
mode="constant",
value=0.0,
)
if (
tenVer.shape[3]
== self.objScratch["levelshape" + str(intRow)][3] + 1
):
tenVer = torch.nn.functional.pad(
input=tenVer,
pad=[0, -1, 0, 0],
mode="constant",
value=0.0,
)
# end
tenIns[intRow] = tenIns[intRow] + tenVer
# end
# end
intRow -= 1
# end
return tenIns
# end
# end
##########################################################
class Network(torch.nn.Module):
def __init__(self):
super().__init__()
self.intEncdec = [1, 1]
self.intChannels = [32, 64, 128, 256, 512]
self.objScratch = {}
self.netInput = torch.nn.Conv2d(
in_channels=3,
out_channels=int(round(0.5 * self.intChannels[0])),
kernel_size=3,
stride=1,
padding=1,
padding_mode="zeros",
)
self.netEncode = torch.nn.Sequential(
*(
[
Encode(
[0] * len(self.intChannels),
self.intChannels,
"prelu(0.25)-conv(3)-prelu(0.25)-conv(3)+skip",
"prelu(0.25)-sconv(3)-prelu(0.25)-conv(3)",
self.objScratch,
)
]
+ [
Encode(
self.intChannels,
self.intChannels,
"prelu(0.25)-conv(3)-prelu(0.25)-conv(3)+skip",
"prelu(0.25)-sconv(3)-prelu(0.25)-conv(3)",
self.objScratch,
)
for intEncdec in range(1, self.intEncdec[0])
]
)
)
self.netDecode = torch.nn.Sequential(
*(
[
Decode(
[0] + self.intChannels[1:],
[0] + self.intChannels[1:],
"prelu(0.25)-conv(3)-prelu(0.25)-conv(3)+skip",
"prelu(0.25)-up(bilinear)-conv(3)-prelu(0.25)-conv(3)",
self.objScratch,
)
for intEncdec in range(0, self.intEncdec[1])
]
)
)
self.netVerone = Basic(
"up(bilinear)-conv(3)-prelu(0.25)-conv(3)",
[self.intChannels[1], self.intChannels[1], 51],
)
self.netVertwo = Basic(
"up(bilinear)-conv(3)-prelu(0.25)-conv(3)",
[self.intChannels[1], self.intChannels[1], 51],
)
self.netHorone = Basic(
"up(bilinear)-conv(3)-prelu(0.25)-conv(3)",
[self.intChannels[1], self.intChannels[1], 51],
)
self.netHortwo = Basic(
"up(bilinear)-conv(3)-prelu(0.25)-conv(3)",
[self.intChannels[1], self.intChannels[1], 51],
)
# self.load_state_dict(torch.hub.load_state_dict_from_url(url='http://content.sniklaus.com/resepconv/network-' + arguments_strModel + '.pytorch', file_name='resepconv-' + arguments_strModel))
# end
def forward(self, x1, x2):
# padding if needed
intWidth = x1.shape[3]
intHeight = x1.shape[2]
intPadr = (2 - (intWidth % 2)) % 2
intPadb = (2 - (intHeight % 2)) % 2
tenOne = torch.nn.functional.pad(
input=x1, pad=[0, intPadr, 0, intPadb], mode="replicate"
)
tenTwo = torch.nn.functional.pad(
input=x2, pad=[0, intPadr, 0, intPadb], mode="replicate"
)
####
tenSeq = [tenOne, tenTwo]
with torch.set_grad_enabled(False):
tenStack = torch.stack(tenSeq, 1)
tenMean = (
tenStack.view(tenStack.shape[0], -1)
.mean(1, True)
.view(tenStack.shape[0], 1, 1, 1)
)
tenStd = (
tenStack.view(tenStack.shape[0], -1)
.std(1, True)
.view(tenStack.shape[0], 1, 1, 1)
)
tenSeq = [
(tenFrame - tenMean) / (tenStd + 0.0000001) for tenFrame in tenSeq
]
tenSeq = [tenFrame.detach() for tenFrame in tenSeq]
# end
tenOut = self.netDecode(
self.netEncode(
[torch.cat([self.netInput(tenSeq[0]), self.netInput(tenSeq[1])], 1)]
+ ([0.0] * (len(self.intChannels) - 1))
)
)[1]
tenOne = torch.nn.functional.pad(
input=tenOne,
pad=[
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
],
mode="replicate",
)
tenTwo = torch.nn.functional.pad(
input=tenTwo,
pad=[
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
int(math.floor(0.5 * 51)),
],
mode="replicate",
)
tenOne = torch.cat(
[
tenOne,
tenOne.new_ones([tenOne.shape[0], 1, tenOne.shape[2], tenOne.shape[3]]),
],
1,
).detach()
tenTwo = torch.cat(
[
tenTwo,
tenTwo.new_ones([tenTwo.shape[0], 1, tenTwo.shape[2], tenTwo.shape[3]]),
],
1,
).detach()
tenVerone = self.netVerone(tenOut)
tenVertwo = self.netVertwo(tenOut)
tenHorone = self.netHorone(tenOut)
tenHortwo = self.netHortwo(tenOut)
tenOut = sepconv_func.apply(tenOne, tenVerone, tenHorone) + sepconv_func.apply(
tenTwo, tenVertwo, tenHortwo
)
tenNormalize = tenOut[:, -1:, :, :]
tenNormalize[tenNormalize.abs() < 0.01] = 1.0
tenOut = tenOut[:, :-1, :, :] / tenNormalize
# crop if needed
return tenOut[:, :, :intHeight, :intWidth]
# end
# end
netNetwork = None
##########################################################
def estimate(tenOne, tenTwo):
global netNetwork
if netNetwork is None:
netNetwork = Network().to(get_torch_device()).eval()
# end
assert tenOne.shape[1] == tenTwo.shape[1]
assert tenOne.shape[2] == tenTwo.shape[2]
intWidth = tenOne.shape[2]
intHeight = tenOne.shape[1]
assert (
intWidth <= 1280
) # while our approach works with larger images, we do not recommend it unless you are aware of the implications
assert (
intHeight <= 720
) # while our approach works with larger images, we do not recommend it unless you are aware of the implications
tenPreprocessedOne = tenOne.to(get_torch_device()).view(1, 3, intHeight, intWidth)
tenPreprocessedTwo = tenTwo.to(get_torch_device()).view(1, 3, intHeight, intWidth)
intPadr = (2 - (intWidth % 2)) % 2
intPadb = (2 - (intHeight % 2)) % 2
tenPreprocessedOne = torch.nn.functional.pad(
input=tenPreprocessedOne, pad=[0, intPadr, 0, intPadb], mode="replicate"
)
tenPreprocessedTwo = torch.nn.functional.pad(
input=tenPreprocessedTwo, pad=[0, intPadr, 0, intPadb], mode="replicate"
)
return netNetwork([tenPreprocessedOne, tenPreprocessedTwo])[
0, :, :intHeight, :intWidth
].cpu()
# end
@@ -1,100 +0,0 @@
import torch
from comfy.model_management import get_torch_device, soft_empty_cache
import numpy as np
import typing
from vfi_utils import InterpolationStateList, load_file_from_github_release, preprocess_frames, postprocess_frames, assert_batch_size
import pathlib
import warnings
import gc
MODEL_TYPE = pathlib.Path(__file__).parent.name
device = get_torch_device()
class STMFNet_VFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (["stmfnet.pth"], ),
"frames": ("IMAGE", ),
"clear_cache_after_n_frames": ("INT", {"default": 10, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 2, "max": 2}), #TODO: Implement recursively invoking interpolator for multi-frame interpolation
"duplicate_first_last_frames": ("BOOLEAN", {"default": False})
},
"optional": {
"optional_interpolation_states": ("INTERPOLATION_STATES", )
}
}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "vfi"
CATEGORY = "ComfyUI-Frame-Interpolation/VFI"
#Reference: https://github.com/danier97/ST-MFNet/blob/main/interpolate_yuv.py#L93
def vfi(
self,
ckpt_name: typing.AnyStr,
frames: torch.Tensor,
clear_cache_after_n_frames = 10,
multiplier: typing.SupportsInt = 2,
duplicate_first_last_frames: bool = False,
optional_interpolation_states: InterpolationStateList = None,
**kwargs
):
from .stmfnet_arch import STMFNet_Model
if multiplier != 2:
warnings.warn("Currently, ST-MFNet only supports 2x interpolation. The process will continue but please set multiplier=2 afterward")
assert_batch_size(frames, batch_size=4, vfi_name="ST-MFNet")
interpolation_states = optional_interpolation_states
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
model = STMFNet_Model()
model.load_state_dict(torch.load(model_path))
model = model.eval().to(device)
frames = preprocess_frames(frames)
number_of_frames_processed_since_last_cleared_cuda_cache = 0
output_frames = []
for frame_itr in range(len(frames) - 3):
#Does skipping frame i+1 make sanse in this case?
if interpolation_states is not None and interpolation_states.is_frame_skipped(frame_itr) and interpolation_states.is_frame_skipped(frame_itr + 1):
continue
#Ensure that input frames are in fp32 - the same dtype as model
frame0, frame1, frame2, frame3 = (
frames[frame_itr:frame_itr+1].float(),
frames[frame_itr+1:frame_itr+2].float(),
frames[frame_itr+2:frame_itr+3].float(),
frames[frame_itr+3:frame_itr+4].float()
)
new_frame = model(frame0.to(device), frame1.to(device), frame2.to(device), frame3.to(device)).detach().cpu()
number_of_frames_processed_since_last_cleared_cuda_cache += 2
if frame_itr == 0:
output_frames.append(frame0)
if duplicate_first_last_frames:
output_frames.append(frame0) # repeat the first frame
output_frames.append(frame1)
output_frames.append(new_frame)
output_frames.append(frame2)
if frame_itr == len(frames) - 4:
output_frames.append(frame3)
if duplicate_first_last_frames:
output_frames.append(frame3) # repeat the last frame
# 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()
dtype = torch.float32
output_frames = [frame.cpu().to(dtype=dtype) 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), )
File diff suppressed because it is too large Load Diff
@@ -1,115 +0,0 @@
import argparse
import torch
import torch.nn as nn
import torch.nn.functional as F
import einops
from torch.utils.data import DataLoader
import pathlib
from vfi_utils import load_file_from_github_release, preprocess_frames, postprocess_frames, InterpolationStateList
import typing
from comfy.model_management import get_torch_device
CKPT_CONFIGS = {
"XVFInet_X4K1000FPS_exp1_latest.pt": {
"module_scale_factor": 4,
"S_trn": 3,
"S_tst": 5
},
"XVFInet_Vimeo_exp1_latest.pt": {
"module_scale_factor": 2,
"S_trn": 1,
"S_tst": 1
}
}
class XVFI_Inference(nn.Module):
def __init__(self, model_path, model_config) -> None:
super(XVFI_Inference, self).__init__()
from .xvfi_arch import XVFInet, weights_init
model_config = model_config
args = argparse.Namespace(
gpu=get_torch_device(),
nf=64,
**model_config,
img_ch=3,
)
self.model = XVFInet(args).apply(weights_init).to(get_torch_device())
self.model.load_state_dict(torch.load(model_path, map_location=get_torch_device())["state_dict_Model"])
def forward(self, I0, I1, timestep):
#"Real" inference is called "test_custom" in the original repo
#https://github.com/JihyongOh/XVFI/blob/main/utils.py#L434
#https://github.com/JihyongOh/XVFI/blob/main/main.py#L336
x = torch.stack([I0, I1], dim=0)
x = einops.rearrange(x, "t b c h w -> b c t h w")
return self.model(x, timestep, is_training=False)
MODEL_TYPE = pathlib.Path(__file__).parent.name
class XVFI:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ckpt_name": (list(CKPT_CONFIGS.keys()), ),
"frames": ("IMAGE", ),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 100}),
"multipler": ("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,
batch_size: typing.SupportsInt = 1,
multipler: typing.SupportsInt = 2,
optional_interpolation_states: InterpolationStateList = None
):
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
ckpt_config = CKPT_CONFIGS[ckpt_name]
global model
model = XVFI_Inference(model_path, ckpt_config)
frames = preprocess_frames(frames)
#https://github.com/JihyongOh/XVFI/blob/main/main.py#L314
divide = 2 ** (ckpt_config["S_tst"]) * ckpt_config["module_scale_factor"] * 4
B, C, H, W = frames.size()
H_padding = (divide - H % divide) % divide
W_padding = (divide - W % divide) % divide
if H_padding != 0 or W_padding != 0:
frames = F.pad(frames, (0, W_padding, 0, H_padding), "constant")
frame_dict = {
str(i): frames[i].unsqueeze(0) for i in range(frames.shape[0])
}
if optional_interpolation_states is None:
interpolation_states = [True] * (frames.shape[0] - 1)
else:
interpolation_states = optional_interpolation_states
enabled_former_idxs = [i for i, state in enumerate(interpolation_states) if state]
former_idxs_loader = DataLoader(enabled_former_idxs, batch_size=batch_size)
for former_idxs_batch in former_idxs_loader:
for middle_i in range(1, multipler):
_middle_frames = model(
frames[former_idxs_batch],
frames[former_idxs_batch + 1],
timestep=torch.tensor([middle_i/multipler]).repeat(len(former_idxs_batch)).unsqueeze(1).to(get_torch_device())
)
for i, former_idx in enumerate(former_idxs_batch):
frame_dict[f'{former_idx}.{middle_i}'] = _middle_frames[i].unsqueeze(0)
out_frames = torch.cat([frame_dict[key] for key in sorted(frame_dict.keys())], dim=0)[:, :, :H, :W]
return (postprocess_frames(out_frames), )
@@ -1,506 +0,0 @@
import functools, random
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.autograd import Variable
import numpy as np
from torch.nn import init
from comfy.model_management import get_torch_device
class XVFInet(nn.Module):
def __init__(self, args):
super(XVFInet, self).__init__()
self.args = args
self.device = get_torch_device()
self.nf = args.nf
self.scale = args.module_scale_factor
self.vfinet = VFInet(args)
self.lrelu = nn.ReLU()
self.in_channels = 3
self.channel_converter = nn.Sequential(
nn.Conv3d(self.in_channels, self.nf, [1, 3, 3], [1, 1, 1], [0, 1, 1]),
nn.ReLU())
self.rec_ext_ds_module = [self.channel_converter]
self.rec_ext_ds = nn.Conv3d(self.nf, self.nf, [1, 3, 3], [1, 2, 2], [0, 1, 1])
for _ in range(int(np.log2(self.scale))):
self.rec_ext_ds_module.append(self.rec_ext_ds)
self.rec_ext_ds_module.append(nn.ReLU())
self.rec_ext_ds_module.append(nn.Conv3d(self.nf, self.nf, [1, 3, 3], 1, [0, 1, 1]))
self.rec_ext_ds_module.append(RResBlock2D_3D(args, T_reduce_flag=False))
self.rec_ext_ds_module = nn.Sequential(*self.rec_ext_ds_module)
self.rec_ctx_ds = nn.Conv3d(self.nf, self.nf, [1, 3, 3], [1, 2, 2], [0, 1, 1])
print("The lowest scale depth for training (S_trn): ", self.args.S_trn)
print("The lowest scale depth for test (S_tst): ", self.args.S_tst)
def forward(self, x, t_value, is_training=True):
'''
x shape : [B,C,T,H,W]
t_value shape : [B,1] ###############
'''
B, C, T, H, W = x.size()
B2, C2 = t_value.size()
assert C2 == 1, "t_value shape is [B,]"
assert T % 2 == 0, "T must be an even number"
t_value = t_value.view(B, 1, 1, 1)
flow_l = None
feat_x = self.rec_ext_ds_module(x)
feat_x_list = [feat_x]
self.lowest_depth_level = self.args.S_trn if is_training else self.args.S_tst
for level in range(1, self.lowest_depth_level+1):
feat_x = self.rec_ctx_ds(feat_x)
feat_x_list.append(feat_x)
if is_training:
out_l_list = []
flow_refine_l_list = []
out_l, flow_l, flow_refine_l = self.vfinet(x, feat_x_list[self.args.S_trn], flow_l, t_value, level=self.args.S_trn, is_training=True)
out_l_list.append(out_l)
flow_refine_l_list.append(flow_refine_l)
for level in range(self.args.S_trn-1, 0, -1): ## self.args.S_trn, self.args.S_trn-1, ..., 1. level 0 is not included
out_l, flow_l = self.vfinet(x, feat_x_list[level], flow_l, t_value, level=level, is_training=True)
out_l_list.append(out_l)
out_l, flow_l, flow_refine_l, occ_0_l0 = self.vfinet(x, feat_x_list[0], flow_l, t_value, level=0, is_training=True)
out_l_list.append(out_l)
flow_refine_l_list.append(flow_refine_l)
return out_l_list[::-1], flow_refine_l_list[::-1], occ_0_l0, torch.mean(x, dim=2) # out_l_list should be reversed. [out_l0, out_l1, ...]
else: # Testing
for level in range(self.args.S_tst, 0, -1): ## self.args.S_tst, self.args.S_tst-1, ..., 1. level 0 is not included
flow_l = self.vfinet(x, feat_x_list[level], flow_l, t_value, level=level, is_training=False)
out_l = self.vfinet(x, feat_x_list[0], flow_l, t_value, level=0, is_training=False)
return out_l
class VFInet(nn.Module):
def __init__(self, args):
super(VFInet, self).__init__()
self.args = args
self.device = get_torch_device()
self.nf = args.nf
self.scale = args.module_scale_factor
self.in_channels = 3
self.conv_flow_bottom = nn.Sequential(
nn.Conv2d(2*self.nf, 2*self.nf, [4,4], 2, [1,1]),
nn.ReLU(),
nn.Conv2d(2*self.nf, 4*self.nf, [4,4], 2, [1,1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(4 * self.nf, 2 * self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(2 * self.nf, self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.Conv2d(self.nf, 6, [3,3], 1, [1,1]),
)
self.conv_flow1 = nn.Conv2d(2*self.nf, self.nf, [3, 3], 1, [1, 1])
self.conv_flow2 = nn.Sequential(
nn.Conv2d(2*self.nf + 4, 2 * self.nf, [4, 4], 2, [1, 1]),
nn.ReLU(),
nn.Conv2d(2 * self.nf, 4 * self.nf, [4, 4], 2, [1, 1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(4 * self.nf, 2 * self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(2 * self.nf, self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.Conv2d(self.nf, 6, [3, 3], 1, [1, 1]),
)
self.conv_flow3 = nn.Sequential(
nn.Conv2d(4 + self.nf * 4, self.nf, [1, 1], 1, [0, 0]),
nn.ReLU(),
nn.Conv2d(self.nf, 2 * self.nf, [4, 4], 2, [1, 1]),
nn.ReLU(),
nn.Conv2d(2 * self.nf, 4 * self.nf, [4, 4], 2, [1, 1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(4 * self.nf, 2 * self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.UpsamplingNearest2d(scale_factor=2),
nn.Conv2d(2 * self.nf, self.nf, [3, 3], 1, [1, 1]),
nn.ReLU(),
nn.Conv2d(self.nf, 4, [3, 3], 1, [1, 1]),
)
self.refine_unet = RefineUNet(args)
self.lrelu = nn.ReLU()
def forward(self, x, feat_x, flow_l_prev, t_value, level, is_training):
'''
x shape : [B,C,T,H,W]
t_value shape : [B,1] ###############
'''
B, C, T, H, W = x.size()
assert T % 2 == 0, "T must be an even number"
####################### For a single level
l = 2 ** level
x_l = x.permute(0,2,1,3,4)
x_l = x_l.contiguous().view(B * T, C, H, W)
if level == 0:
pass
else:
x_l = F.interpolate(x_l, scale_factor=(1.0 / l, 1.0 / l), mode='bicubic', align_corners=False)
'''
Down pixel-shuffle
'''
x_l = x_l.view(B, T, C, H//l, W//l)
x_l = x_l.permute(0,2,1,3,4)
B, C, T, H, W = x_l.size()
## Feature extraction
feat0_l = feat_x[:,:,0,:,:]
feat1_l = feat_x[:,:,1,:,:]
## Flow estimation
if flow_l_prev is None:
flow_l_tmp = self.conv_flow_bottom(torch.cat((feat0_l, feat1_l), dim=1))
flow_l = flow_l_tmp[:,:4,:,:]
else:
up_flow_l_prev = 2.0*F.interpolate(flow_l_prev.detach(), scale_factor=(2,2), mode='bilinear', align_corners=False)
warped_feat1_l = self.bwarp(feat1_l, up_flow_l_prev[:,:2,:,:])
warped_feat0_l = self.bwarp(feat0_l, up_flow_l_prev[:,2:,:,:])
flow_l_tmp = self.conv_flow2(torch.cat([self.conv_flow1(torch.cat([feat0_l, warped_feat1_l],dim=1)), self.conv_flow1(torch.cat([feat1_l, warped_feat0_l],dim=1)), up_flow_l_prev],dim=1))
flow_l = flow_l_tmp[:,:4,:,:] + up_flow_l_prev
if not is_training and level!=0:
return flow_l
flow_01_l = flow_l[:,:2,:,:]
flow_10_l = flow_l[:,2:,:,:]
z_01_l = torch.sigmoid(flow_l_tmp[:,4:5,:,:])
z_10_l = torch.sigmoid(flow_l_tmp[:,5:6,:,:])
## Complementary Flow Reversal (CFR)
flow_forward, norm0_l = self.z_fwarp(flow_01_l, t_value * flow_01_l, z_01_l) ## Actually, F (t) -> (t+1). Translation only. Not normalized yet
flow_backward, norm1_l = self.z_fwarp(flow_10_l, (1-t_value) * flow_10_l, z_10_l) ## Actually, F (1-t) -> (-t). Translation only. Not normalized yet
flow_t0_l = -(1-t_value) * ((t_value)*flow_forward) + (t_value) * ((t_value)*flow_backward) # The numerator of Eq.(1) in the paper.
flow_t1_l = (1-t_value) * ((1-t_value)*flow_forward) - (t_value) * ((1-t_value)*flow_backward) # The numerator of Eq.(2) in the paper.
norm_l = (1-t_value)*norm0_l + t_value*norm1_l
mask_ = (norm_l.detach() > 0).type(norm_l.type())
flow_t0_l = (1-mask_) * flow_t0_l + mask_ * (flow_t0_l.clone() / (norm_l.clone() + (1-mask_))) # Divide the numerator with denominator in Eq.(1)
flow_t1_l = (1-mask_) * flow_t1_l + mask_ * (flow_t1_l.clone() / (norm_l.clone() + (1-mask_))) # Divide the numerator with denominator in Eq.(2)
## Feature warping
warped0_l = self.bwarp(feat0_l, flow_t0_l)
warped1_l = self.bwarp(feat1_l, flow_t1_l)
## Flow refinement
flow_refine_l = torch.cat([feat0_l, warped0_l, warped1_l, feat1_l, flow_t0_l, flow_t1_l], dim=1)
flow_refine_l = self.conv_flow3(flow_refine_l) + torch.cat([flow_t0_l, flow_t1_l], dim=1)
flow_t0_l = flow_refine_l[:, :2, :, :]
flow_t1_l = flow_refine_l[:, 2:4, :, :]
warped0_l = self.bwarp(feat0_l, flow_t0_l)
warped1_l = self.bwarp(feat1_l, flow_t1_l)
## Flow upscale
flow_t0_l = self.scale * F.interpolate(flow_t0_l, scale_factor=(self.scale, self.scale), mode='bilinear',align_corners=False)
flow_t1_l = self.scale * F.interpolate(flow_t1_l, scale_factor=(self.scale, self.scale), mode='bilinear',align_corners=False)
## Image warping and blending
warped_img0_l = self.bwarp(x_l[:,:,0,:,:], flow_t0_l)
warped_img1_l = self.bwarp(x_l[:,:,1,:,:], flow_t1_l)
refine_out = self.refine_unet(torch.cat([F.pixel_shuffle(torch.cat([feat0_l, feat1_l, warped0_l, warped1_l],dim=1), self.scale), x_l[:,:,0,:,:], x_l[:,:,1,:,:], warped_img0_l, warped_img1_l, flow_t0_l, flow_t1_l],dim=1))
occ_0_l = torch.sigmoid(refine_out[:, 0:1, :, :])
occ_1_l = 1-occ_0_l
out_l = (1-t_value)*occ_0_l*warped_img0_l + t_value*occ_1_l*warped_img1_l
out_l = out_l / ( (1-t_value)*occ_0_l + t_value*occ_1_l ) + refine_out[:, 1:4, :, :]
if not is_training and level==0:
return out_l
if is_training:
if flow_l_prev is None:
# if level == self.args.S_trn:
return out_l, flow_l, flow_refine_l[:, 0:4, :, :]
elif level != 0:
return out_l, flow_l
else: # level==0
return out_l, flow_l, flow_refine_l[:, 0:4, :, :], occ_0_l
def bwarp(self, x, flo):
'''
x: [B, C, H, W] (im2)
flo: [B, 2, H, W] flow
'''
B, C, H, W = x.size()
# mesh grid
xx = torch.arange(0, W).view(1, 1, 1, W).expand(B, 1, H, W)
yy = torch.arange(0, H).view(1, 1, H, 1).expand(B, 1, H, W)
grid = torch.cat((xx, yy), 1).float()
grid = grid.to(self.device)
vgrid = torch.autograd.Variable(grid) + flo
# scale grid to [-1,1]
vgrid[:, 0, :, :] = 2.0 * vgrid[:, 0, :, :].clone() / max(W - 1, 1) - 1.0
vgrid[:, 1, :, :] = 2.0 * vgrid[:, 1, :, :].clone() / max(H - 1, 1) - 1.0
vgrid = vgrid.permute(0, 2, 3, 1) # [B,H,W,2]
output = nn.functional.grid_sample(x, vgrid, align_corners=True)
mask = torch.autograd.Variable(torch.ones(x.size())).to(self.device)
mask = nn.functional.grid_sample(mask, vgrid, align_corners=True)
# mask[mask<0.9999] = 0
# mask[mask>0] = 1
mask = mask.masked_fill_(mask < 0.999, 0)
mask = mask.masked_fill_(mask > 0, 1)
return output * mask
def fwarp(self, img, flo):
"""
-img: image (N, C, H, W)
-flo: optical flow (N, 2, H, W)
elements of flo is in [0, H] and [0, W] for dx, dy
https://github.com/lyh-18/EQVI/blob/EQVI-master/models/forward_warp_gaussian.py
"""
# (x1, y1) (x1, y2)
# +---------------+
# | |
# | o(x, y) |
# | |
# | |
# | |
# | |
# +---------------+
# (x2, y1) (x2, y2)
N, C, _, _ = img.size()
# translate start-point optical flow to end-point optical flow
y = flo[:, 0:1:, :]
x = flo[:, 1:2, :, :]
x = x.repeat(1, C, 1, 1)
y = y.repeat(1, C, 1, 1)
# Four point of square (x1, y1), (x1, y2), (x2, y1), (y2, y2)
x1 = torch.floor(x)
x2 = x1 + 1
y1 = torch.floor(y)
y2 = y1 + 1
# firstly, get gaussian weights
w11, w12, w21, w22 = self.get_gaussian_weights(x, y, x1, x2, y1, y2)
# secondly, sample each weighted corner
img11, o11 = self.sample_one(img, x1, y1, w11)
img12, o12 = self.sample_one(img, x1, y2, w12)
img21, o21 = self.sample_one(img, x2, y1, w21)
img22, o22 = self.sample_one(img, x2, y2, w22)
imgw = img11 + img12 + img21 + img22
o = o11 + o12 + o21 + o22
return imgw, o
def z_fwarp(self, img, flo, z):
"""
-img: image (N, C, H, W)
-flo: optical flow (N, 2, H, W)
elements of flo is in [0, H] and [0, W] for dx, dy
modified from https://github.com/lyh-18/EQVI/blob/EQVI-master/models/forward_warp_gaussian.py
"""
# (x1, y1) (x1, y2)
# +---------------+
# | |
# | o(x, y) |
# | |
# | |
# | |
# | |
# +---------------+
# (x2, y1) (x2, y2)
N, C, _, _ = img.size()
# translate start-point optical flow to end-point optical flow
y = flo[:, 0:1:, :]
x = flo[:, 1:2, :, :]
x = x.repeat(1, C, 1, 1)
y = y.repeat(1, C, 1, 1)
# Four point of square (x1, y1), (x1, y2), (x2, y1), (y2, y2)
x1 = torch.floor(x)
x2 = x1 + 1
y1 = torch.floor(y)
y2 = y1 + 1
# firstly, get gaussian weights
w11, w12, w21, w22 = self.get_gaussian_weights(x, y, x1, x2, y1, y2, z+1e-5)
# secondly, sample each weighted corner
img11, o11 = self.sample_one(img, x1, y1, w11)
img12, o12 = self.sample_one(img, x1, y2, w12)
img21, o21 = self.sample_one(img, x2, y1, w21)
img22, o22 = self.sample_one(img, x2, y2, w22)
imgw = img11 + img12 + img21 + img22
o = o11 + o12 + o21 + o22
return imgw, o
def get_gaussian_weights(self, x, y, x1, x2, y1, y2, z=1.0):
# z 0.0 ~ 1.0
w11 = z * torch.exp(-((x - x1) ** 2 + (y - y1) ** 2))
w12 = z * torch.exp(-((x - x1) ** 2 + (y - y2) ** 2))
w21 = z * torch.exp(-((x - x2) ** 2 + (y - y1) ** 2))
w22 = z * torch.exp(-((x - x2) ** 2 + (y - y2) ** 2))
return w11, w12, w21, w22
def sample_one(self, img, shiftx, shifty, weight):
"""
Input:
-img (N, C, H, W)
-shiftx, shifty (N, c, H, W)
"""
N, C, H, W = img.size()
# flatten all (all restored as Tensors)
flat_shiftx = shiftx.view(-1)
flat_shifty = shifty.view(-1)
flat_basex = torch.arange(0, H, requires_grad=False).view(-1, 1)[None, None].to(self.device).long().repeat(N, C,1,W).view(-1)
flat_basey = torch.arange(0, W, requires_grad=False).view(1, -1)[None, None].to(self.device).long().repeat(N, C,H,1).view(-1)
flat_weight = weight.view(-1)
flat_img = img.contiguous().view(-1)
# The corresponding positions in I1
idxn = torch.arange(0, N, requires_grad=False).view(N, 1, 1, 1).to(self.device).long().repeat(1, C, H, W).view(-1)
idxc = torch.arange(0, C, requires_grad=False).view(1, C, 1, 1).to(self.device).long().repeat(N, 1, H, W).view(-1)
idxx = flat_shiftx.long() + flat_basex
idxy = flat_shifty.long() + flat_basey
# recording the inside part the shifted
mask = idxx.ge(0) & idxx.lt(H) & idxy.ge(0) & idxy.lt(W)
# Mask off points out of boundaries
ids = (idxn * C * H * W + idxc * H * W + idxx * W + idxy)
ids_mask = torch.masked_select(ids, mask).clone().to(self.device)
# Note here! accmulate fla must be true for proper bp
img_warp = torch.zeros([N * C * H * W, ]).to(self.device)
img_warp.put_(ids_mask, torch.masked_select(flat_img * flat_weight, mask), accumulate=True)
one_warp = torch.zeros([N * C * H * W, ]).to(self.device)
one_warp.put_(ids_mask, torch.masked_select(flat_weight, mask), accumulate=True)
return img_warp.view(N, C, H, W), one_warp.view(N, C, H, W)
class RefineUNet(nn.Module):
def __init__(self, args):
super(RefineUNet, self).__init__()
self.args = args
self.scale = args.module_scale_factor
self.nf = args.nf
self.conv1 = nn.Conv2d(self.nf, self.nf, [3,3], 1, [1,1])
self.conv2 = nn.Conv2d(self.nf, self.nf, [3,3], 1, [1,1])
self.lrelu = nn.ReLU()
self.NN = nn.UpsamplingNearest2d(scale_factor=2)
self.enc1 = nn.Conv2d((4*self.nf)//self.scale//self.scale + 4*args.img_ch + 4, self.nf, [4, 4], 2, [1, 1])
self.enc2 = nn.Conv2d(self.nf, 2*self.nf, [4, 4], 2, [1, 1])
self.enc3 = nn.Conv2d(2*self.nf, 4*self.nf, [4, 4], 2, [1, 1])
self.dec0 = nn.Conv2d(4*self.nf, 4*self.nf, [3, 3], 1, [1, 1])
self.dec1 = nn.Conv2d(4*self.nf + 2*self.nf, 2*self.nf, [3, 3], 1, [1, 1]) ## input concatenated with enc2
self.dec2 = nn.Conv2d(2*self.nf + self.nf, self.nf, [3, 3], 1, [1, 1]) ## input concatenated with enc1
self.dec3 = nn.Conv2d(self.nf, 1+args.img_ch, [3, 3], 1, [1, 1]) ## input added with warped image
def forward(self, concat):
enc1 = self.lrelu(self.enc1(concat))
enc2 = self.lrelu(self.enc2(enc1))
out = self.lrelu(self.enc3(enc2))
out = self.lrelu(self.dec0(out))
out = self.NN(out)
out = torch.cat((out,enc2),dim=1)
out = self.lrelu(self.dec1(out))
out = self.NN(out)
out = torch.cat((out,enc1),dim=1)
out = self.lrelu(self.dec2(out))
out = self.NN(out)
out = self.dec3(out)
return out
class ResBlock2D_3D(nn.Module):
## Shape of input [B,C,T,H,W]
## Shape of output [B,C,T,H,W]
def __init__(self, args):
super(ResBlock2D_3D, self).__init__()
self.args = args
self.nf = args.nf
self.conv3x3_1 = nn.Conv3d(self.nf, self.nf, [1,3,3], 1, [0,1,1])
self.conv3x3_2 = nn.Conv3d(self.nf, self.nf, [1,3,3], 1, [0,1,1])
self.lrelu = nn.ReLU()
def forward(self, x):
'''
x shape : [B,C,T,H,W]
'''
B, C, T, H, W = x.size()
out = self.conv3x3_2(self.lrelu(self.conv3x3_1(x)))
return x + out
class RResBlock2D_3D(nn.Module):
def __init__(self, args, T_reduce_flag=False):
super(RResBlock2D_3D, self).__init__()
self.args = args
self.nf = args.nf
self.T_reduce_flag = T_reduce_flag
self.resblock1 = ResBlock2D_3D(self.args)
self.resblock2 = ResBlock2D_3D(self.args)
if T_reduce_flag:
self.reduceT_conv = nn.Conv3d(self.nf, self.nf, [3,1,1], 1, [0,0,0])
def forward(self, x):
'''
x shape : [B,C,T,H,W]
'''
out = self.resblock1(x)
out = self.resblock2(out)
if self.T_reduce_flag:
return self.reduceT_conv(out + x)
else:
return out + x
def weights_init(m):
classname = m.__class__.__name__
if (classname.find('Conv2d') != -1) or (classname.find('Conv3d') != -1):
init.xavier_normal_(m.weight)
# init.kaiming_normal_(m.weight, nonlinearity='relu')
if hasattr(m, 'bias') and m.bias is not None:
init.zeros_(m.bias)
@@ -24,7 +24,7 @@ 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 InterpolationStateList():
class InterpolationStateListImport():
def __init__(self, frame_indices: typing.List[int], is_skip_list: bool):
self.frame_indices = frame_indices
@@ -35,7 +35,7 @@ class InterpolationStateList():
return self.is_skip_list and is_frame_in_list or not self.is_skip_list and not is_frame_in_list
class MakeInterpolationStateList:
class MakeInterpolationStateListImport:
@classmethod
def INPUT_TYPES(s):
return {
@@ -52,7 +52,7 @@ class MakeInterpolationStateList:
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 = InterpolationStateList(
interpolation_state_list = InterpolationStateListImport(
frame_indices=frame_indices_list,
is_skip_list=is_skip_list,
)
@@ -127,7 +127,7 @@ def _generic_frame_loop(
multiplier: typing.Union[typing.SupportsInt, typing.List],
return_middle_frame_function,
*return_middle_frame_function_args,
interpolation_states: InterpolationStateList = None,
interpolation_states: InterpolationStateListImport = None,
use_timestep=True,
dtype=torch.float16,
final_logging=True):
@@ -213,7 +213,7 @@ def generic_frame_loop(
multiplier: typing.Union[typing.SupportsInt, typing.List],
return_middle_frame_function,
*return_middle_frame_function_args,
interpolation_states: InterpolationStateList = None,
interpolation_states: InterpolationStateListImport = None,
use_timestep=True,
dtype=torch.float32):
@@ -255,7 +255,7 @@ def generic_frame_loop(
return output_frames
raise NotImplementedError(f"multipiler of {type(multiplier)}")
class FloatToInt:
class FloatToIntImport:
@classmethod
def INPUT_TYPES(s):
return {