| 257 | |
| 258 | @torch.no_grad() |
| 259 | def validate(config, data_loader, model, logger): |
| 260 | criterion = nn.CrossEntropyLoss() |
| 261 | model.eval() |
| 262 | |
| 263 | batch_time = AverageMeter() |
| 264 | loss_meter = AverageMeter() |
| 265 | acc1_meter = AverageMeter() |
| 266 | acc5_meter = AverageMeter() |
| 267 | |
| 268 | end = time.time() |
| 269 | for idx, (images, target) in enumerate(data_loader): |
| 270 | images = images.cuda(non_blocking=True) |
| 271 | target = target.cuda(non_blocking=True) |
| 272 | |
| 273 | # compute output |
| 274 | output, _, _ = model(images) |
| 275 | |
| 276 | # measure accuracy and record loss |
| 277 | loss = criterion(output, target) |
| 278 | acc1, acc5 = accuracy(output, target, topk=(1, 5)) |
| 279 | |
| 280 | acc1 = reduce_tensor(acc1) |
| 281 | acc5 = reduce_tensor(acc5) |
| 282 | loss = reduce_tensor(loss) |
| 283 | |
| 284 | loss_meter.update(loss.item(), target.size(0)) |
| 285 | acc1_meter.update(acc1.item(), target.size(0)) |
| 286 | acc5_meter.update(acc5.item(), target.size(0)) |
| 287 | |
| 288 | # measure elapsed time |
| 289 | batch_time.update(time.time() - end) |
| 290 | end = time.time() |
| 291 | |
| 292 | if (idx + 1) % config.PRINT_FREQ == 0: |
| 293 | memory_used = torch.cuda.max_memory_allocated() / (1024.0 * 1024.0) |
| 294 | logger.info( |
| 295 | f'Test: [{(idx + 1)}/{len(data_loader)}]\t' |
| 296 | f'Time {batch_time.val:.3f} ({batch_time.avg:.3f})\t' |
| 297 | f'Loss {loss_meter.val:.4f} ({loss_meter.avg:.4f})\t' |
| 298 | f'Acc@1 {acc1_meter.val:.3f} ({acc1_meter.avg:.3f})\t' |
| 299 | f'Acc@5 {acc5_meter.val:.3f} ({acc5_meter.avg:.3f})\t' |
| 300 | f'Mem {memory_used:.0f}MB') |
| 301 | logger.info(f' * Acc@1 {acc1_meter.avg:.3f} Acc@5 {acc5_meter.avg:.3f}') |
| 302 | return acc1_meter.avg, acc5_meter.avg, loss_meter.avg |
| 303 | |
| 304 | @torch.no_grad() |
| 305 | def throughput(data_loader, model, logger): |