CoolFace
Apppublic

LeFleur808/Free-View_Expressive_Talking_Head_Video_Editing

sourceHugging Facecc-by-nc-4.0updated 1y agoView on Hugging Face
0likes
fete_model.py297 linesDownload Raw Back to root
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