CoolFace
Modelpublic

Isi99999/Frame_Interpolation_Models

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
2likes
IFNet_HDv3.py116 linesDownload Raw Back to train_log
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