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

Method train_network

talkNet.py:21–49  ·  view source on GitHub ↗
(self, loader, epoch, **kwargs)

Source from the content-addressed store, hash-verified

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()

Callers 1

mainFunction · 0.80

Calls 8

forward_audio_backendMethod · 0.80
forwardMethod · 0.45
__len__Method · 0.45

Tested by

no test coverage detected