MCPcopy Create free account
hub / github.com/deeplearning-wisc/MOOD / validate

Function validate

utils/msdnet_function.py:22–87  ·  view source on GitHub ↗
(val_loader, model, criterion)

Source from the content-addressed store, hash-verified

20 args.num_classes = 1000
21
22def 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]))

Callers 1

main.pyFile · 0.90

Calls 3

updateMethod · 0.95
accuracyFunction · 0.85
AverageMeterClass · 0.70

Tested by

no test coverage detected