MCPcopy Create free account
hub / github.com/TaoRuijie/TalkNet-ASD / talkNet

Class talkNet

talkNet.py:10–94  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8from model.talkNetModel import talkNetModel
9
10class talkNet(nn.Module):
11 def __init__(self, lr = 0.0001, lrDecay = 0.95, **kwargs):
12 super(talkNet, self).__init__()
13 self.model = talkNetModel().cuda()
14 self.lossAV = lossAV().cuda()
15 self.lossA = lossA().cuda()
16 self.lossV = lossV().cuda()
17 self.optim = torch.optim.Adam(self.parameters(), lr = lr)
18 self.scheduler = torch.optim.lr_scheduler.StepLR(self.optim, step_size = 1, gamma=lrDecay)
19 print(time.strftime("%m-%d %H:%M:%S") + " Model para number = %.2f"%(sum(param.numel() for param in self.model.parameters()) / 1024 / 1024))
20
21 def train_network(self, loader, epoch, **kwargs):
22 self.train()
23 self.scheduler.step(epoch - 1)
24 index, top1, loss = 0, 0, 0
25 lr = self.optim.param_groups[0]['lr']
26 for num, (audioFeature, visualFeature, labels) in enumerate(loader, start=1):
27 self.zero_grad()
28 audioEmbed = self.model.forward_audio_frontend(audioFeature[0].cuda()) # feedForward
29 visualEmbed = self.model.forward_visual_frontend(visualFeature[0].cuda())
30 audioEmbed, visualEmbed = self.model.forward_cross_attention(audioEmbed, visualEmbed)
31 outsAV= self.model.forward_audio_visual_backend(audioEmbed, visualEmbed)
32 outsA = self.model.forward_audio_backend(audioEmbed)
33 outsV = self.model.forward_visual_backend(visualEmbed)
34 labels = labels[0].reshape((-1)).cuda() # Loss
35 nlossAV, _, _, prec = self.lossAV.forward(outsAV, labels)
36 nlossA = self.lossA.forward(outsA, labels)
37 nlossV = self.lossV.forward(outsV, labels)
38 nloss = nlossAV + 0.4 * nlossA + 0.4 * nlossV
39 loss += nloss.detach().cpu().numpy()
40 top1 += prec
41 nloss.backward()
42 self.optim.step()
43 index += len(labels)
44 sys.stderr.write(time.strftime("%m-%d %H:%M:%S") + \
45 " [%2d] Lr: %5f, Training: %.2f%%, " %(epoch, lr, 100 * (num / loader.__len__())) + \
46 " Loss: %.5f, ACC: %2.2f%% \r" %(loss/(num), 100 * (top1/index)))
47 sys.stderr.flush()
48 sys.stdout.write("\n")
49 return loss/num, lr
50
51 def evaluate_network(self, loader, evalCsvSave, evalOrig, **kwargs):
52 self.eval()
53 predScores = []
54 for audioFeature, visualFeature, labels in tqdm.tqdm(loader):
55 with torch.no_grad():
56 audioEmbed = self.model.forward_audio_frontend(audioFeature[0].cuda())
57 visualEmbed = self.model.forward_visual_frontend(visualFeature[0].cuda())
58 audioEmbed, visualEmbed = self.model.forward_cross_attention(audioEmbed, visualEmbed)
59 outsAV= self.model.forward_audio_visual_backend(audioEmbed, visualEmbed)
60 labels = labels[0].reshape((-1)).cuda()
61 _, predScore, _, _ = self.lossAV.forward(outsAV, labels)
62 predScore = predScore[:,1].detach().cpu().numpy()
63 predScores.extend(predScore)
64 evalLines = open(evalOrig).read().splitlines()[1:]
65 labels = []
66 labels = pandas.Series( ['SPEAKING_AUDIBLE' for line in evalLines])
67 scores = pandas.Series(predScores)

Callers 2

evaluate_networkFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected