syncnet
| 235 | |
| 236 | |
| 237 | class SyncNet(nn.Module): |
| 238 | """ |
| 239 | syncnet |
| 240 | """ |
| 241 | |
| 242 | def __init__(self, in_channel_image, in_channel_audio, out_dim): |
| 243 | super(SyncNet, self).__init__() |
| 244 | self.in_channel_image = in_channel_image |
| 245 | self.in_channel_audio = in_channel_audio |
| 246 | self.out_dim = out_dim |
| 247 | self.face_encoder = FaceEncoder(in_channel_image, out_dim) |
| 248 | self.audio_encoder = AudioEncoder(in_channel_audio, out_dim) |
| 249 | self.merge_encoder = nn.Sequential( |
| 250 | nn.Conv2d(out_dim * 2, out_dim, kernel_size=3, padding=1), |
| 251 | nn.LeakyReLU(0.2), |
| 252 | nn.Conv2d(out_dim, 1, kernel_size=3, padding=1), |
| 253 | ) |
| 254 | |
| 255 | def forward(self, image, audio): |
| 256 | image_embedding = self.face_encoder(image) |
| 257 | audio_embedding = ( |
| 258 | self.audio_encoder(audio) |
| 259 | .unsqueeze(2) |
| 260 | .unsqueeze(3) |
| 261 | .repeat(1, 1, image_embedding.size(2), image_embedding.size(3)) |
| 262 | ) |
| 263 | concat_embedding = torch.cat([image_embedding, audio_embedding], 1) |
| 264 | out_score = self.merge_encoder(concat_embedding) |
| 265 | return out_score |
| 266 | |
| 267 | |
| 268 | class SyncNetPerception(nn.Module): |