FIxing classes
This commit is contained in:
+3
-2
@@ -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
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
-1857
File diff suppressed because it is too large
Load Diff
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user