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

Method evaluate_network

talkNet.py:51–76  ·  view source on GitHub ↗
(self, loader, evalCsvSave, evalOrig, **kwargs)

Source from the content-addressed store, hash-verified

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)
68 evalRes = pandas.read_csv(evalOrig)
69 evalRes['score'] = scores
70 evalRes['label'] = labels
71 evalRes.drop(['label_id'], axis=1,inplace=True)
72 evalRes.drop(['instance_id'], axis=1,inplace=True)
73 evalRes.to_csv(evalCsvSave, index=False)
74 cmd = "python -O utils/get_ava_active_speaker_performance.py -g %s -p %s "%(evalOrig, evalCsvSave)
75 mAP = float(str(subprocess.run(cmd, shell=True, capture_output =True).stdout).split(' ')[2][:5])
76 return mAP
77
78 def saveParameters(self, path):
79 torch.save(self.state_dict(), path)

Callers 1

mainFunction · 0.80

Calls 5

forwardMethod · 0.45

Tested by

no test coverage detected