CoolFace
Apppublic

LexDF/CogVideoX-5B-Space

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
IFNet_HDv3.py139 linesDownload Raw Back to rife
1import torch2import torch.nn as nn3import torch.nn.functional as F4from .warplayer import warp5 6device = torch.device("cuda" if torch.cuda.is_available() else "cpu")7 8 9def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):10    return nn.Sequential(11        nn.Conv2d(12            in_planes,13            out_planes,14            kernel_size=kernel_size,15            stride=stride,16            padding=padding,17            dilation=dilation,18            bias=True,19        ),20        nn.PReLU(out_planes),21    )22 23 24def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):25    return nn.Sequential(26        nn.Conv2d(27            in_planes,28            out_planes,29            kernel_size=kernel_size,30            stride=stride,31            padding=padding,32            dilation=dilation,33            bias=False,34        ),35        nn.BatchNorm2d(out_planes),36        nn.PReLU(out_planes),37    )38 39 40class IFBlock(nn.Module):41    def __init__(self, in_planes, c=64):42        super(IFBlock, self).__init__()43        self.conv0 = nn.Sequential(44            conv(in_planes, c // 2, 3, 2, 1),45            conv(c // 2, c, 3, 2, 1),46        )47        self.convblock0 = nn.Sequential(conv(c, c), conv(c, c))48        self.convblock1 = nn.Sequential(conv(c, c), conv(c, c))49        self.convblock2 = nn.Sequential(conv(c, c), conv(c, c))50        self.convblock3 = nn.Sequential(conv(c, c), conv(c, c))51        self.conv1 = nn.Sequential(52            nn.ConvTranspose2d(c, c // 2, 4, 2, 1),53            nn.PReLU(c // 2),54            nn.ConvTranspose2d(c // 2, 4, 4, 2, 1),55        )56        self.conv2 = nn.Sequential(57            nn.ConvTranspose2d(c, c // 2, 4, 2, 1),58            nn.PReLU(c // 2),59            nn.ConvTranspose2d(c // 2, 1, 4, 2, 1),60        )61 62    def forward(self, x, flow, scale=1):63        x = F.interpolate(64            x, scale_factor=1.0 / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False65        )66        flow = (67            F.interpolate(68                flow, scale_factor=1.0 / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False69            )70            * 1.071            / scale72        )73        feat = self.conv0(torch.cat((x, flow), 1))74        feat = self.convblock0(feat) + feat75        feat = self.convblock1(feat) + feat76        feat = self.convblock2(feat) + feat77        feat = self.convblock3(feat) + feat78        flow = self.conv1(feat)79        mask = self.conv2(feat)80        flow = (81            F.interpolate(flow, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)82            * scale83        )84        mask = F.interpolate(85            mask, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False86        )87        return flow, mask88 89 90class IFNet(nn.Module):91    def __init__(self):92        super(IFNet, self).__init__()93        self.block0 = IFBlock(7 + 4, c=90)94        self.block1 = IFBlock(7 + 4, c=90)95        self.block2 = IFBlock(7 + 4, c=90)96        self.block_tea = IFBlock(10 + 4, c=90)97        # self.contextnet = Contextnet()98        # self.unet = Unet()99 100    def forward(self, x, scale_list=[4, 2, 1], training=False):101        if training == False:102            channel = x.shape[1] // 2103            img0 = x[:, :channel]104            img1 = x[:, channel:]105        flow_list = []106        merged = []107        mask_list = []108        warped_img0 = img0109        warped_img1 = img1110        flow = (x[:, :4]).detach() * 0111        mask = (x[:, :1]).detach() * 0112        loss_cons = 0113        block = [self.block0, self.block1, self.block2]114        for i in range(3):115            f0, m0 = block[i](torch.cat((warped_img0[:, :3], warped_img1[:, :3], mask), 1), flow, scale=scale_list[i])116            f1, m1 = block[i](117                torch.cat((warped_img1[:, :3], warped_img0[:, :3], -mask), 1),118                torch.cat((flow[:, 2:4], flow[:, :2]), 1),119                scale=scale_list[i],120            )121            flow = flow + (f0 + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 2122            mask = mask + (m0 + (-m1)) / 2123            mask_list.append(mask)124            flow_list.append(flow)125            warped_img0 = warp(img0, flow[:, :2])126            warped_img1 = warp(img1, flow[:, 2:4])127            merged.append((warped_img0, warped_img1))128        """129        c0 = self.contextnet(img0, flow[:, :2])130        c1 = self.contextnet(img1, flow[:, 2:4])131        tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)132        res = tmp[:, 1:4] * 2 - 1133        """134        for i in range(3):135            mask_list[i] = torch.sigmoid(mask_list[i])136            merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])137            # merged[i] = torch.clamp(merged[i] + res, 0, 1)138        return flow_list, mask_list[2], merged139