| 109 | |
| 110 | |
| 111 | def train(train_queue, model, criterion, optimizer): |
| 112 | objs = utils.AvgrageMeter() |
| 113 | top1 = utils.AvgrageMeter() |
| 114 | top5 = utils.AvgrageMeter() |
| 115 | model.train() |
| 116 | |
| 117 | for step, (input, target) in enumerate(train_queue): |
| 118 | input = Variable(input).cuda() |
| 119 | target = Variable(target).cuda(non_blocking=True) |
| 120 | |
| 121 | optimizer.zero_grad() |
| 122 | logits, logits_aux = model(input) |
| 123 | loss = criterion(logits, target) |
| 124 | if args.auxiliary: |
| 125 | loss_aux = criterion(logits_aux, target) |
| 126 | loss += args.auxiliary_weight*loss_aux |
| 127 | loss.backward() |
| 128 | nn.utils.clip_grad_norm(model.parameters(), args.grad_clip) |
| 129 | optimizer.step() |
| 130 | |
| 131 | prec1, prec5 = utils.accuracy(logits, target, topk=(1, 5)) |
| 132 | n = input.size(0) |
| 133 | objs.update(loss.data[0], n) |
| 134 | top1.update(prec1.data[0], n) |
| 135 | top5.update(prec5.data[0], n) |
| 136 | |
| 137 | if step % args.report_freq == 0: |
| 138 | logging.info('train %03d %e %f %f', step, objs.avg, top1.avg, top5.avg) |
| 139 | |
| 140 | return top1.avg, objs.avg |
| 141 | |
| 142 | |
| 143 | def infer(valid_queue, model, criterion): |