bleckhert/Free-View_Expressive_Talking_Head_Video_Editing
0
1import torch2from torch import nn3from torch.nn import functional as F4 5 6class Conv2d(nn.Module):7 def __init__(self, cin, cout, kernel_size, stride, padding, *args, **kwargs):8 super().__init__(*args, **kwargs)9 self.conv_block = nn.Sequential(nn.Conv2d(cin, cout, kernel_size, stride, padding), nn.BatchNorm2d(cout))10 self.act = nn.ReLU()11 12 def forward(self, x):13 out = self.conv_block(x)14 return self.act(out)15 16 17class Conv2d_res(nn.Module):18 # TensorRT does not support 'if' statement, thus we create independent Conv2d_res for residual block19 def __init__(self, cin, cout, kernel_size, stride, padding, *args, **kwargs):20 super().__init__(*args, **kwargs)21 self.conv_block = nn.Sequential(nn.Conv2d(cin, cout, kernel_size, stride, padding), nn.BatchNorm2d(cout))22 self.act = nn.ReLU()23 24 def forward(self, x):25 out = self.conv_block(x)26 out += x27 return self.act(out)28 29 30class Conv2dTranspose(nn.Module):31 def __init__(self, cin, cout, kernel_size, stride, padding, output_padding=0, *args, **kwargs):32 super().__init__(*args, **kwargs)33 self.conv_block = nn.Sequential(34 nn.ConvTranspose2d(cin, cout, kernel_size, stride, padding, output_padding),35 nn.BatchNorm2d(cout),36 )37 self.act = nn.ReLU()38 39 def forward(self, x):40 out = self.conv_block(x)41 return self.act(out)42 43 44class FETE_model(nn.Module):45 def __init__(self):46 super(FETE_model, self).__init__()47 48 self.face_encoder_blocks = nn.ModuleList(49 [50 nn.Sequential(Conv2d(6, 16, kernel_size=7, stride=2, padding=3)), # 256,256 -> 128,12851 nn.Sequential(52 Conv2d(16, 32, kernel_size=3, stride=2, padding=1), # 64,6453 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),54 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),55 ),56 nn.Sequential(57 Conv2d(32, 64, kernel_size=3, stride=2, padding=1), # 32,3258 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),59 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),60 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),61 ),62 nn.Sequential(63 Conv2d(64, 128, kernel_size=3, stride=2, padding=1), # 16,1664 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),65 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),66 ),67 nn.Sequential(68 Conv2d(128, 256, kernel_size=3, stride=2, padding=1), # 8,869 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),70 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),71 ),72 nn.Sequential(73 Conv2d(256, 512, kernel_size=3, stride=2, padding=1), # 4,474 Conv2d_res(512, 512, kernel_size=3, stride=1, padding=1),75 ),76 nn.Sequential(77 Conv2d(512, 512, kernel_size=3, stride=2, padding=0), # 1, 178 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),79 ),80 ]81 )82 83 self.audio_encoder = nn.Sequential(84 Conv2d(1, 32, kernel_size=3, stride=1, padding=1),85 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),86 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),87 Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1),88 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),89 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),90 Conv2d(64, 128, kernel_size=3, stride=3, padding=1),91 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),92 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),93 Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1),94 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),95 Conv2d(256, 512, kernel_size=3, stride=1, padding=0),96 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),97 )98 99 self.pose_encoder = nn.Sequential(100 Conv2d(1, 32, kernel_size=3, stride=1, padding=1),101 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),102 Conv2d(32, 64, kernel_size=3, stride=(1, 2), padding=1),103 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),104 Conv2d(64, 128, kernel_size=3, stride=1, padding=1),105 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),106 Conv2d(128, 256, kernel_size=3, stride=(1, 2), padding=1),107 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),108 Conv2d(256, 512, kernel_size=3, stride=2, padding=0),109 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),110 )111 112 self.emotion_encoder = nn.Sequential(113 Conv2d(1, 32, kernel_size=7, stride=1, padding=1),114 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),115 Conv2d(32, 64, kernel_size=3, stride=(1, 2), padding=1),116 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),117 Conv2d(64, 128, kernel_size=3, stride=1, padding=1),118 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),119 Conv2d(128, 256, kernel_size=3, stride=(1, 2), padding=1),120 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),121 Conv2d(256, 512, kernel_size=3, stride=2, padding=0),122 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),123 )124 125 self.blink_encoder = nn.Sequential(126 Conv2d(1, 32, kernel_size=3, stride=1, padding=1),127 Conv2d_res(32, 32, kernel_size=3, stride=1, padding=1),128 Conv2d(32, 64, kernel_size=3, stride=(1, 2), padding=1),129 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),130 Conv2d(64, 128, kernel_size=3, stride=(1, 2), padding=1),131 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),132 Conv2d(128, 256, kernel_size=3, stride=(1, 2), padding=1),133 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),134 Conv2d(256, 512, kernel_size=1, stride=(1, 2), padding=0),135 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),136 )137 138 self.face_decoder_blocks = nn.ModuleList(139 [140 nn.Sequential(141 Conv2d(2048, 512, kernel_size=1, stride=1, padding=0),142 ),143 nn.Sequential(144 Conv2dTranspose(1024, 512, kernel_size=4, stride=1, padding=0), # 4,4145 Conv2d_res(512, 512, kernel_size=3, stride=1, padding=1),146 ),147 nn.Sequential(148 Conv2dTranspose(1024, 512, kernel_size=3, stride=2, padding=1, output_padding=1),149 Conv2d_res(512, 512, kernel_size=3, stride=1, padding=1),150 Conv2d_res(512, 512, kernel_size=3, stride=1, padding=1), # 8,8151 Self_Attention(512, 512),152 ),153 nn.Sequential(154 Conv2dTranspose(768, 384, kernel_size=3, stride=2, padding=1, output_padding=1),155 Conv2d_res(384, 384, kernel_size=3, stride=1, padding=1),156 Conv2d_res(384, 384, kernel_size=3, stride=1, padding=1), # 16, 16157 Self_Attention(384, 384),158 ),159 nn.Sequential(160 Conv2dTranspose(512, 256, kernel_size=3, stride=2, padding=1, output_padding=1),161 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1),162 Conv2d_res(256, 256, kernel_size=3, stride=1, padding=1), # 32, 32163 Self_Attention(256, 256),164 ),165 nn.Sequential(166 Conv2dTranspose(320, 128, kernel_size=3, stride=2, padding=1, output_padding=1),167 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),168 Conv2d_res(128, 128, kernel_size=3, stride=1, padding=1),169 ), # 64, 64170 nn.Sequential(171 Conv2dTranspose(160, 64, kernel_size=3, stride=2, padding=1, output_padding=1),172 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),173 Conv2d_res(64, 64, kernel_size=3, stride=1, padding=1),174 ),175 ]176 ) # 128,128177 178 # self.output_block = nn.Sequential(Conv2d(80, 32, kernel_size=3, stride=1, padding=1),179 # nn.Conv2d(32, 3, kernel_size=1, stride=1, padding=0),180 # nn.Sigmoid())181 182 self.output_block = nn.Sequential(183 Conv2dTranspose(80, 32, kernel_size=3, stride=2, padding=1, output_padding=1),184 nn.Conv2d(32, 3, kernel_size=1, stride=1, padding=0),185 nn.Sigmoid(),186 )187 188 def forward(189 self,190 face_sequences,191 audio_sequences,192 pose_sequences,193 emotion_sequences,194 blink_sequences,195 ):196 # audio_sequences = (B, T, 1, 80, 16)197 B = audio_sequences.size(0)198 199 # disabled for inference200 # input_dim_size = len(face_sequences.size())201 # if input_dim_size > 4:202 # audio_sequences = torch.cat([audio_sequences[:, i] for i in range(audio_sequences.size(1))], dim=0)203 # pose_sequences = torch.cat([pose_sequences[:, i] for i in range(pose_sequences.size(1))], dim=0)204 # emotion_sequences = torch.cat([emotion_sequences[:, i] for i in range(emotion_sequences.size(1))], dim=0)205 # blink_sequences = torch.cat([blink_sequences[:, i] for i in range(blink_sequences.size(1))], dim=0)206 # face_sequences = torch.cat([face_sequences[:, :, i] for i in range(face_sequences.size(2))], dim=0)207 # print(audio_sequences.size(), face_sequences.size(), pose_sequences.size(), emotion_sequences.size())208 209 audio_embedding = self.audio_encoder(audio_sequences) # B, 512, 1, 1210 pose_embedding = self.pose_encoder(pose_sequences) # B, 512, 1, 1211 emotion_embedding = self.emotion_encoder(emotion_sequences) # B, 512, 1, 1212 blink_embedding = self.blink_encoder(blink_sequences) # B, 512, 1, 1213 inputs_embedding = torch.cat((audio_embedding, pose_embedding, emotion_embedding, blink_embedding), dim=1) # B, 1536, 1, 1214 # print(audio_embedding.size(), pose_embedding.size(), emotion_embedding.size(), inputs_embedding.size())215 216 feats = []217 x = face_sequences218 for f in self.face_encoder_blocks:219 x = f(x)220 # print(x.shape)221 feats.append(x)222 223 x = inputs_embedding224 for f in self.face_decoder_blocks:225 x = f(x)226 # print(x.shape)227 228 # try:229 x = torch.cat((x, feats[-1]), dim=1)230 # except Exception as e:231 # print(x.size())232 # print(feats[-1].size())233 # raise e234 feats.pop()235 236 x = self.output_block(x)237 238 # if input_dim_size > 4:239 # x = torch.split(x, B, dim=0) # [(B, C, H, W)]240 # outputs = torch.stack(x, dim=2) # (B, C, T, H, W)241 242 # else:243 outputs = x244 245 return outputs246 247 248class Self_Attention(nn.Module):249 """250 Source-Reference Attention Layer251 """252 253 def __init__(self, in_planes_s, in_planes_r):254 """255 Parameters256 ----------257 in_planes_s: int258 Number of input source feature vector channels.259 in_planes_r: int260 Number of input reference feature vector channels.261 """262 super(Self_Attention, self).__init__()263 self.query_conv = nn.Conv2d(in_channels=in_planes_s, out_channels=in_planes_s // 8, kernel_size=1)264 self.key_conv = nn.Conv2d(in_channels=in_planes_r, out_channels=in_planes_r // 8, kernel_size=1)265 self.value_conv = nn.Conv2d(in_channels=in_planes_r, out_channels=in_planes_r, kernel_size=1)266 self.gamma = nn.Parameter(torch.zeros(1))267 self.softmax = nn.Softmax(dim=-1)268 269 def forward(self, source):270 source = source.float() if isinstance(source, torch.cuda.HalfTensor) else source271 reference = source272 """273 Parameters274 ----------275 source : torch.Tensor276 Source feature maps (B x Cs x Ts x Hs x Ws)277 reference : torch.Tensor278 Reference feature maps (B x Cr x Tr x Hr x Wr )279 Returns :280 torch.Tensor281 Source-reference attention value added to the input source features282 torch.Tensor283 Attention map (B x Ns x Nt) (Ns=Ts*Hs*Ws, Nr=Tr*Hr*Wr)284 """285 s_batchsize, sC, sH, sW = source.size()286 r_batchsize, rC, rH, rW = reference.size()287 288 proj_query = self.query_conv(source).view(s_batchsize, -1, sH * sW).permute(0, 2, 1)289 proj_key = self.key_conv(reference).view(r_batchsize, -1, rW * rH)290 energy = torch.bmm(proj_query, proj_key)291 attention = self.softmax(energy)292 proj_value = self.value_conv(reference).view(r_batchsize, -1, rH * rW)293 out = torch.bmm(proj_value, attention.permute(0, 2, 1))294 out = out.view(s_batchsize, sC, sH, sW)295 out = self.gamma * out + source296 return out.half() if isinstance(source, torch.cuda.FloatTensor) else out297 