MCPcopy Create free account
hub / github.com/10Ring/LAA-Net / validate

Function validate

lib/core_function.py:172–275  ·  view source on GitHub ↗
(cfg, model, critetion, epoch, data_loader, logger, writer, devices, valIters, metrics_base='combine')

Source from the content-addressed store, hash-verified

170
171
172def validate(cfg, model, critetion, epoch, data_loader, logger, writer, devices, valIters, metrics_base='combine'):
173 calculate_acc = get_acc_mesure_func(metrics_base)
174 batch_time = AverageMeter()
175 data_time = AverageMeter()
176 losses = AverageMeter()
177 acc = AverageMeter()
178
179 #Switch to test mode
180 model.eval()
181 data_loader = tqdm(data_loader, dynamic_ncols=True)
182 start = time.time()
183 with torch.no_grad():
184 for i, batch_data in enumerate(data_loader):
185 inputs, labels, targets, heatmaps, cstency_heatmaps, offsets = get_batch_data(batch_data)
186 inputs = inputs.to(devices, non_blocking=True, dtype=torch.float64).cuda()
187 #Measuring data loading time
188 data_time.update(time.time() - start)
189
190 outputs = model(inputs)
191 if isinstance(outputs, list):
192 outputs = outputs[0]
193 #In case outputs contain a dict key
194 if isinstance(outputs, dict):
195 outputs_hm = outputs['hm']
196 outputs_cls = outputs['cls']
197 outputs_offset = outputs['offset'] if 'offset' in outputs.keys() else None
198 outputs_cstency = outputs['cstency'] if 'cstency' in outputs.keys() else None
199
200 if 'Combined' in cfg.TRAIN.loss.type:
201 labels = labels.cuda().to(non_blocking=True, dtype=torch.float64)
202 # labels = labels.cuda().to(non_blocking=True).long()
203
204 if offsets is not None:
205 offsets = offsets.cuda().to(non_blocking=True, dtype=torch.float64)
206
207 if cstency_heatmaps is not None:
208 cstency_heatmaps = cstency_heatmaps.cuda().to(non_blocking=True, dtype=torch.float64)
209
210 if cfg.TRAIN.loss.type != 'CombinedHeatmapBinaryLoss':
211 heatmaps = heatmaps.cuda().to(non_blocking=True, dtype=torch.float64)
212 else:
213 heatmaps = targets.cuda().to(non_blocking=True, dtype=torch.float64)
214
215 loss_ = critetion(outputs_hm, heatmaps, outputs_cls.sigmoid(), labels,
216 offset_preds=outputs_offset,
217 offset_gts=offsets,
218 cstency_preds=outputs_cstency,
219 cstency_gts=cstency_heatmaps)
220 loss = loss_['hm']
221 if 'cls' in loss_.keys():
222 loss += loss_['cls']
223 if 'dst_hm_cls' in loss_.keys():
224 loss += loss_['dst_hm_cls']
225 if 'offset' in loss_.keys():
226 loss += loss_['offset']
227 if 'cstency' in loss_.keys():
228 loss += loss_['cstency']
229 else:

Callers 1

train.pyFile · 0.90

Calls 7

updateMethod · 0.95
get_acc_mesure_funcFunction · 0.90
debugging_panelFunction · 0.90
board_writingFunction · 0.90
AverageMeterClass · 0.85
get_batch_dataFunction · 0.85
epochInforMethod · 0.80

Tested by

no test coverage detected