LexDF/CogVideoX-5B-Space
0
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 