(val_loader, model, criterion)
| 20 | args.num_classes = 1000 |
| 21 | |
| 22 | def validate(val_loader, model, criterion): |
| 23 | print(args.nBlocks) |
| 24 | batch_time = AverageMeter() |
| 25 | losses = AverageMeter() |
| 26 | data_time = AverageMeter() |
| 27 | top1, top5 = [], [] |
| 28 | for i in range(args.nBlocks): |
| 29 | top1.append(AverageMeter()) |
| 30 | top5.append(AverageMeter()) |
| 31 | |
| 32 | # switch to evaluate mode |
| 33 | model.eval() |
| 34 | |
| 35 | end = time.time() |
| 36 | with torch.no_grad(): |
| 37 | for i, (input, target) in enumerate(val_loader): |
| 38 | target = target.cuda(non_blocking=True) |
| 39 | input = input.cuda() |
| 40 | |
| 41 | input_var = torch.autograd.Variable(input) |
| 42 | target_var = torch.autograd.Variable(target) |
| 43 | |
| 44 | data_time.update(time.time() - end) |
| 45 | |
| 46 | # compute output |
| 47 | output, _ = model(input_var) |
| 48 | if not isinstance(output, list): |
| 49 | output = [output] |
| 50 | |
| 51 | loss = 0.0 |
| 52 | for j in range(len(output)): |
| 53 | loss += criterion(output[j], target_var) |
| 54 | |
| 55 | # measure error and record loss |
| 56 | losses.update(loss.item(), input.size(0)) |
| 57 | |
| 58 | for j in range(len(output)): |
| 59 | err1, err5 = accuracy(output[j].data, target, topk=(1, 5)) |
| 60 | top1[j].update(err1.item(), input.size(0)) |
| 61 | top5[j].update(err5.item(), input.size(0)) |
| 62 | |
| 63 | # measure elapsed time |
| 64 | batch_time.update(time.time() - end) |
| 65 | end = time.time() |
| 66 | |
| 67 | if i % args.print_freq == 0: |
| 68 | print('Epoch: [{0}/{1}]\t' |
| 69 | 'Time {batch_time.avg:.3f}\t' |
| 70 | 'Data {data_time.avg:.3f}\t' |
| 71 | 'Loss {loss.val:.4f}\t' |
| 72 | 'Err@1 {top1.val:.4f}\t' |
| 73 | 'Err@5 {top5.val:.4f}'.format( |
| 74 | i + 1, len(val_loader), |
| 75 | batch_time=batch_time, data_time=data_time, |
| 76 | loss=losses, top1=top1[-1], top5=top5[-1])) |
| 77 | # break |
| 78 | for j in range(args.nBlocks): |
| 79 | print(' * Err@1 {top1.avg:.3f} Err@5 {top5.avg:.3f}'.format(top1=top1[j], top5=top5[j])) |
no test coverage detected