(cfg, model, critetion, epoch, data_loader, logger, writer, devices, valIters, metrics_base='combine')
| 170 | |
| 171 | |
| 172 | def 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: |
no test coverage detected