MCPcopy Create free account
hub / github.com/OUCMachineLearning/OUCML / train

Function train

AutoML/darts-master/cnn/train.py:111–140  ·  view source on GitHub ↗
(train_queue, model, criterion, optimizer)

Source from the content-addressed store, hash-verified

109
110
111def 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
143def infer(valid_queue, model, criterion):

Callers 1

mainFunction · 0.70

Calls 3

updateMethod · 0.95
trainMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected