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

Class SyncNet

models/old/Syncnet_halfBN.py:238–262  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

236
237
238class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected