Isi99999/Frame_Interpolation_Models
2
1import torch2import torch.nn as nn3import torch.nn.functional as F4from model.warplayer import warp5 6device = torch.device("cuda" if torch.cuda.is_available() else "cpu")7 8def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):9 return nn.Sequential(10 nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,11 padding=padding, dilation=dilation, bias=True), 12 nn.PReLU(out_planes)13 )14 15def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):16 return nn.Sequential(17 nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,18 padding=padding, dilation=dilation, bias=False),19 nn.BatchNorm2d(out_planes),20 nn.PReLU(out_planes)21 )22 23class IFBlock(nn.Module):24 def __init__(self, in_planes, c=64):25 super(IFBlock, self).__init__()26 self.conv0 = nn.Sequential(27 conv(in_planes, c//2, 3, 2, 1),28 conv(c//2, c, 3, 2, 1),29 )30 self.convblock0 = nn.Sequential(31 conv(c, c),32 conv(c, c)33 )34 self.convblock1 = nn.Sequential(35 conv(c, c),36 conv(c, c)37 )38 self.convblock2 = nn.Sequential(39 conv(c, c),40 conv(c, c)41 )42 self.convblock3 = nn.Sequential(43 conv(c, c),44 conv(c, c)45 )46 self.conv1 = nn.Sequential(47 nn.ConvTranspose2d(c, c//2, 4, 2, 1),48 nn.PReLU(c//2),49 nn.ConvTranspose2d(c//2, 4, 4, 2, 1),50 )51 self.conv2 = nn.Sequential(52 nn.ConvTranspose2d(c, c//2, 4, 2, 1),53 nn.PReLU(c//2),54 nn.ConvTranspose2d(c//2, 1, 4, 2, 1),55 )56 57 def forward(self, x, flow, scale=1):58 x = F.interpolate(x, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)59 flow = F.interpolate(flow, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 1. / scale60 feat = self.conv0(torch.cat((x, flow), 1))61 feat = self.convblock0(feat) + feat62 feat = self.convblock1(feat) + feat63 feat = self.convblock2(feat) + feat64 feat = self.convblock3(feat) + feat 65 flow = self.conv1(feat)66 mask = self.conv2(feat)67 flow = F.interpolate(flow, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * scale68 mask = F.interpolate(mask, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)69 return flow, mask70 71class IFNet(nn.Module):72 def __init__(self):73 super(IFNet, self).__init__()74 self.block0 = IFBlock(7+4, c=90)75 self.block1 = IFBlock(7+4, c=90)76 self.block2 = IFBlock(7+4, c=90)77 self.block_tea = IFBlock(10+4, c=90)78 # self.contextnet = Contextnet()79 # self.unet = Unet()80 81 def forward(self, x, scale_list=[4, 2, 1], training=False):82 if training == False:83 channel = x.shape[1] // 284 img0 = x[:, :channel]85 img1 = x[:, channel:]86 flow_list = []87 merged = []88 mask_list = []89 warped_img0 = img090 warped_img1 = img191 flow = (x[:, :4]).detach() * 092 mask = (x[:, :1]).detach() * 093 loss_cons = 094 block = [self.block0, self.block1, self.block2]95 for i in range(3):96 f0, m0 = block[i](torch.cat((warped_img0[:, :3], warped_img1[:, :3], mask), 1), flow, scale=scale_list[i])97 f1, m1 = block[i](torch.cat((warped_img1[:, :3], warped_img0[:, :3], -mask), 1), torch.cat((flow[:, 2:4], flow[:, :2]), 1), scale=scale_list[i])98 flow = flow + (f0 + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 299 mask = mask + (m0 + (-m1)) / 2100 mask_list.append(mask)101 flow_list.append(flow)102 warped_img0 = warp(img0, flow[:, :2])103 warped_img1 = warp(img1, flow[:, 2:4])104 merged.append((warped_img0, warped_img1))105 '''106 c0 = self.contextnet(img0, flow[:, :2])107 c1 = self.contextnet(img1, flow[:, 2:4])108 tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)109 res = tmp[:, 1:4] * 2 - 1110 '''111 for i in range(3):112 mask_list[i] = torch.sigmoid(mask_list[i])113 merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])114 # merged[i] = torch.clamp(merged[i] + res, 0, 1) 115 return flow_list, mask_list[2], merged116 