neuralleap/CogVideoX-5B-API-V2
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(nn.Module):61 def __init__(self):62 super(IFNet, self).__init__()63 self.block0 = IFBlock(6, c=240)64 self.block1 = IFBlock(13 + 4, c=150)65 self.block2 = IFBlock(13 + 4, c=90)66 self.block_tea = IFBlock(16 + 4, c=90)67 self.contextnet = Contextnet()68 self.unet = Unet()69 70 def forward(self, x, scale=[4, 2, 1], timestep=0.5):71 img0 = x[:, :3]72 img1 = x[:, 3:6]73 gt = x[:, 6:] # In inference time, gt is None74 flow_list = []75 merged = []76 mask_list = []77 warped_img0 = img078 warped_img1 = img179 flow = None80 loss_distill = 081 stu = [self.block0, self.block1, self.block2]82 for i in range(3):83 if flow != None:84 flow_d, mask_d = stu[i](85 torch.cat((img0, img1, warped_img0, warped_img1, mask), 1), flow, scale=scale[i]86 )87 flow = flow + flow_d88 mask = mask + mask_d89 else:90 flow, mask = stu[i](torch.cat((img0, img1), 1), None, scale=scale[i])91 mask_list.append(torch.sigmoid(mask))92 flow_list.append(flow)93 warped_img0 = warp(img0, flow[:, :2])94 warped_img1 = warp(img1, flow[:, 2:4])95 merged_student = (warped_img0, warped_img1)96 merged.append(merged_student)97 if gt.shape[1] == 3:98 flow_d, mask_d = self.block_tea(99 torch.cat((img0, img1, warped_img0, warped_img1, mask, gt), 1), flow, scale=1100 )101 flow_teacher = flow + flow_d102 warped_img0_teacher = warp(img0, flow_teacher[:, :2])103 warped_img1_teacher = warp(img1, flow_teacher[:, 2:4])104 mask_teacher = torch.sigmoid(mask + mask_d)105 merged_teacher = warped_img0_teacher * mask_teacher + warped_img1_teacher * (1 - mask_teacher)106 else:107 flow_teacher = None108 merged_teacher = None109 for i in range(3):110 merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])111 if gt.shape[1] == 3:112 loss_mask = (113 ((merged[i] - gt).abs().mean(1, True) > (merged_teacher - gt).abs().mean(1, True) + 0.01)114 .float()115 .detach()116 )117 loss_distill += (((flow_teacher.detach() - flow_list[i]) ** 2).mean(1, True) ** 0.5 * loss_mask).mean()118 c0 = self.contextnet(img0, flow[:, :2])119 c1 = self.contextnet(img1, flow[:, 2:4])120 tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)121 res = tmp[:, :3] * 2 - 1122 merged[2] = torch.clamp(merged[2] + res, 0, 1)123 return flow_list, mask_list[2], merged, flow_teacher, merged_teacher, loss_distill124 