(self)
| 436 | return losses.avg, top1.avg, top5.avg |
| 437 | |
| 438 | def validate(self): |
| 439 | losses = AverageMeter() |
| 440 | top1 = AverageMeter() |
| 441 | top5 = AverageMeter() |
| 442 | # switch to evaluate mode |
| 443 | self.network.eval() |
| 444 | for i, (inputs, targets) in enumerate(self.test_loader): |
| 445 | targets = targets.to(self.gpu) |
| 446 | inputs = inputs.to(self.gpu) |
| 447 | with torch.no_grad(): |
| 448 | # compute output |
| 449 | outputs = self.network(inputs) |
| 450 | loss = self.criterion(outputs, targets) |
| 451 | |
| 452 | # measure accuracy and record loss |
| 453 | prec1, prec5 = accuracy(outputs.data, targets, topk=(1, 5)) |
| 454 | losses.update(loss.item(), inputs.size(0)) |
| 455 | top1.update(prec1[0], inputs.size(0)) |
| 456 | top5.update(prec5[0], inputs.size(0)) |
| 457 | return losses.avg, top1.avg, top5.avg |
| 458 | |
| 459 | |
| 460 | class AverageMeter: |
no test coverage detected