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

Class SyncNet

models/Syncnet.py:237–265  ·  view source on GitHub ↗

syncnet

Source from the content-addressed store, hash-verified

235
236
237class 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
268class SyncNetPerception(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected