| 8 | from model.talkNetModel import talkNetModel |
| 9 | |
| 10 | class 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) |
no outgoing calls
no test coverage detected