MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / val

Function val

CV/Pytorch_classification/CRNN_PyTorch/train.py:127–171  ·  view source on GitHub ↗
(net, dataset, criterion, max_iter=100)

Source from the content-addressed store, hash-verified

125
126
127def val(net, dataset, criterion, max_iter=100):
128 print('Start Val')
129
130 for p in crnn.parameters():
131 p.requires_grad = False
132
133 net.eval()
134 data_loader = torch.utils.data.DataLoader(
135 dataset, shuffle=True, batch_size=opt.batchSize, num_workers=int(opt.workers))
136 val_iter = iter(data_loader)
137
138 i = 0
139 n_correct = 0
140 loss_avg = utils.averager()
141
142 max_iter = min(max_iter, len(data_loader))
143 for i in range(max_iter):
144 data = val_iter.next()
145 i += 1
146 cpu_images, cpu_texts = data
147 batch_size = cpu_images.size(0)
148 utils.loadData(image, cpu_images)
149 t, l = converter.encode(cpu_texts)
150 utils.loadData(text, t)
151 utils.loadData(length, l)
152
153 preds = crnn(image)
154 preds_size = Variable(torch.IntTensor([preds.size(0)] * batch_size))
155 cost = criterion(preds, text, preds_size.data, raw=False)
156 loss_avg.add(cost)
157
158 _, preds = preds.max(2)
159 preds = preds.squeeze(2)
160 preds = preds.transpose(1, 0).contiguous().view(-1)
161 sim_preds = converter.decode(preds.data, preds_size.data, raw=False)
162 for pred, target in zip(sim_preds, cpu_texts):
163 if pred == target.lower():
164 n_correct += 1
165
166 raw_preds = converter.decode(preds.data, preds_size.data, raw=True)[:opt.n_test_disp]
167 for raw_pred, pred, gt in zip(raw_preds, sim_preds, cpu_texts):
168 print('%-20s => %-20s, gt: %-20s' % (raw_pred, pred, gt))
169
170 accuracy = n_correct / float(max_iter * opt.batchSize)
171 print('Test loss: %f, accuracy: %f' %(loss_avg.val(), accuracy))
172
173def trainBatch(net, ctiterion, optimizer):
174 data = train_iter.next()

Callers 1

train.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected