MCPcopy Create free account
hub / github.com/Elsaam2y/DINet_optimized / SyncNet

Class SyncNet

models/old/Syncnet_BN.py:253–277  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

251
252
253class SyncNet(nn.Module):
254 def __init__(self, in_channel_image, in_channel_audio, out_dim):
255 super(SyncNet, self).__init__()
256 self.in_channel_image = in_channel_image
257 self.in_channel_audio = in_channel_audio
258 self.out_dim = out_dim
259 self.face_encoder = FaceEncoder(in_channel_image, out_dim)
260 self.audio_encoder = AudioEncoder(in_channel_audio, out_dim)
261 self.merge_encoder = nn.Sequential(
262 nn.Conv2d(out_dim * 2, out_dim, kernel_size=3, padding=1),
263 nn.LeakyReLU(0.2),
264 nn.Conv2d(out_dim, 1, kernel_size=3, padding=1),
265 )
266
267 def forward(self, image, audio):
268 image_embedding = self.face_encoder(image)
269 audio_embedding = (
270 self.audio_encoder(audio)
271 .unsqueeze(2)
272 .unsqueeze(3)
273 .repeat(1, 1, image_embedding.size(2), image_embedding.size(3))
274 )
275 concat_embedding = torch.cat([image_embedding, audio_embedding], 1)
276 out_score = self.merge_encoder(concat_embedding)
277 return out_score

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected