CoolFace
Datasetpublic

H-Liu1997/tango_cached_utils

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes230downloads
syncnet.py67 linesDownload Raw Back to models
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