| 236 | |
| 237 | |
| 238 | class SyncNet(nn.Module): |
| 239 | def __init__(self, in_channel_image, in_channel_audio, out_dim): |
| 240 | super(SyncNet, self).__init__() |
| 241 | self.in_channel_image = in_channel_image |
| 242 | self.in_channel_audio = in_channel_audio |
| 243 | self.out_dim = out_dim |
| 244 | self.face_encoder = FaceEncoder(in_channel_image, out_dim) |
| 245 | self.audio_encoder = AudioEncoder(in_channel_audio, out_dim) |
| 246 | self.merge_encoder = nn.Sequential( |
| 247 | nn.Conv2d(out_dim * 2, out_dim, kernel_size=3, padding=1), |
| 248 | nn.LeakyReLU(0.2), |
| 249 | nn.Conv2d(out_dim, 1, kernel_size=3, padding=1), |
| 250 | ) |
| 251 | |
| 252 | def forward(self, image, audio): |
| 253 | image_embedding = self.face_encoder(image) |
| 254 | audio_embedding = ( |
| 255 | self.audio_encoder(audio) |
| 256 | .unsqueeze(2) |
| 257 | .unsqueeze(3) |
| 258 | .repeat(1, 1, image_embedding.size(2), image_embedding.size(3)) |
| 259 | ) |
| 260 | concat_embedding = torch.cat([image_embedding, audio_embedding], 1) |
| 261 | out_score = self.merge_encoder(concat_embedding) |
| 262 | return out_score |
nothing calls this directly
no outgoing calls
no test coverage detected