H-Liu1997/tango_cached_utils
0230
1import torch2from torch import nn3from torch.nn import functional as F4 5from .conv import Conv2d6 7class SyncNet_color(nn.Module):8 def __init__(self):9 super(SyncNet_color, self).__init__()10 11 self.face_encoder = nn.Sequential(12 Conv2d(15, 32, kernel_size=(7, 7), stride=1, padding=3),13 14 Conv2d(32, 64, kernel_size=5, stride=(1, 2), padding=1),15 Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),16 Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),17 18 Conv2d(64, 128, kernel_size=3, stride=2, padding=1),19 Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),20 Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),21 Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),22 23 Conv2d(128, 256, kernel_size=3, stride=2, padding=1),24 Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),25 Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),26 27 Conv2d(256, 512, kernel_size=3, stride=2, padding=1),28 Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True),29 Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True),30 31 Conv2d(512, 512, kernel_size=3, stride=2, padding=1),32 Conv2d(512, 512, kernel_size=3, stride=1, padding=0),33 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),)34 35 self.audio_encoder = nn.Sequential(36 Conv2d(1, 32, kernel_size=3, stride=1, padding=1),37 Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),38 Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),39 40 Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1),41 Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),42 Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),43 44 Conv2d(64, 128, kernel_size=3, stride=3, padding=1),45 Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),46 Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),47 48 Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1),49 Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),50 Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),51 52 Conv2d(256, 512, kernel_size=3, stride=1, padding=0),53 Conv2d(512, 512, kernel_size=1, stride=1, padding=0),)54 55 def forward(self, audio_sequences, face_sequences): # audio_sequences := (B, dim, T)56 face_embedding = self.face_encoder(face_sequences)57 audio_embedding = self.audio_encoder(audio_sequences)58 59 audio_embedding = audio_embedding.view(audio_embedding.size(0), -1)60 face_embedding = face_embedding.view(face_embedding.size(0), -1)61 62 audio_embedding = F.normalize(audio_embedding, p=2, dim=1)63 face_embedding = F.normalize(face_embedding, p=2, dim=1)64 65 66 return audio_embedding, face_embedding67 