LexDF/CogVideoX-5B-Space
0
1from .refine import *2 3 4def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):5 return nn.Sequential(6 torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1),7 nn.PReLU(out_planes),8 )9 10 11def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):12 return nn.Sequential(13 nn.Conv2d(14 in_planes,15 out_planes,16 kernel_size=kernel_size,17 stride=stride,18 padding=padding,19 dilation=dilation,20 bias=True,21 ),22 nn.PReLU(out_planes),23 )24 25 26class IFBlock(nn.Module):27 def __init__(self, in_planes, c=64):28 super(IFBlock, self).__init__()29 self.conv0 = nn.Sequential(30 conv(in_planes, c // 2, 3, 2, 1),31 conv(c // 2, c, 3, 2, 1),32 )33 self.convblock = nn.Sequential(34 conv(c, c),35 conv(c, c),36 conv(c, c),37 conv(c, c),38 conv(c, c),39 conv(c, c),40 conv(c, c),41 conv(c, c),42 )43 self.lastconv = nn.ConvTranspose2d(c, 5, 4, 2, 1)44 45 def forward(self, x, flow, scale):46 if scale != 1:47 x = F.interpolate(x, scale_factor=1.0 / scale, mode="bilinear", align_corners=False)48 if flow != None:49 flow = F.interpolate(flow, scale_factor=1.0 / scale, mode="bilinear", align_corners=False) * 1.0 / scale50 x = torch.cat((x, flow), 1)51 x = self.conv0(x)52 x = self.convblock(x) + x53 tmp = self.lastconv(x)54 tmp = F.interpolate(tmp, scale_factor=scale * 2, mode="bilinear", align_corners=False)55 flow = tmp[:, :4] * scale * 256 mask = tmp[:, 4:5]57 return flow, mask58 59 60class IFNet_m(nn.Module):61 def __init__(self):62 super(IFNet_m, self).__init__()63 self.block0 = IFBlock(6 + 1, c=240)64 self.block1 = IFBlock(13 + 4 + 1, c=150)65 self.block2 = IFBlock(13 + 4 + 1, c=90)66 self.block_tea = IFBlock(16 + 4 + 1, c=90)67 self.contextnet = Contextnet()68 self.unet = Unet()69 70 def forward(self, x, scale=[4, 2, 1], timestep=0.5, returnflow=False):71 timestep = (x[:, :1].clone() * 0 + 1) * timestep72 img0 = x[:, :3]73 img1 = x[:, 3:6]74 gt = x[:, 6:] # In inference time, gt is None75 flow_list = []76 merged = []77 mask_list = []78 warped_img0 = img079 warped_img1 = img180 flow = None81 loss_distill = 082 stu = [self.block0, self.block1, self.block2]83 for i in range(3):84 if flow != None:85 flow_d, mask_d = stu[i](86 torch.cat((img0, img1, timestep, warped_img0, warped_img1, mask), 1), flow, scale=scale[i]87 )88 flow = flow + flow_d89 mask = mask + mask_d90 else:91 flow, mask = stu[i](torch.cat((img0, img1, timestep), 1), None, scale=scale[i])92 mask_list.append(torch.sigmoid(mask))93 flow_list.append(flow)94 warped_img0 = warp(img0, flow[:, :2])95 warped_img1 = warp(img1, flow[:, 2:4])96 merged_student = (warped_img0, warped_img1)97 merged.append(merged_student)98 if gt.shape[1] == 3:99 flow_d, mask_d = self.block_tea(100 torch.cat((img0, img1, timestep, warped_img0, warped_img1, mask, gt), 1), flow, scale=1101 )102 flow_teacher = flow + flow_d103 warped_img0_teacher = warp(img0, flow_teacher[:, :2])104 warped_img1_teacher = warp(img1, flow_teacher[:, 2:4])105 mask_teacher = torch.sigmoid(mask + mask_d)106 merged_teacher = warped_img0_teacher * mask_teacher + warped_img1_teacher * (1 - mask_teacher)107 else:108 flow_teacher = None109 merged_teacher = None110 for i in range(3):111 merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])112 if gt.shape[1] == 3:113 loss_mask = (114 ((merged[i] - gt).abs().mean(1, True) > (merged_teacher - gt).abs().mean(1, True) + 0.01)115 .float()116 .detach()117 )118 loss_distill += (((flow_teacher.detach() - flow_list[i]) ** 2).mean(1, True) ** 0.5 * loss_mask).mean()119 if returnflow:120 return flow121 else:122 c0 = self.contextnet(img0, flow[:, :2])123 c1 = self.contextnet(img1, flow[:, 2:4])124 tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)125 res = tmp[:, :3] * 2 - 1126 merged[2] = torch.clamp(merged[2] + res, 0, 1)127 return flow_list, mask_list[2], merged, flow_teacher, merged_teacher, loss_distill128 