(model, classifier, dataloader, args)
| 71 | |
| 72 | |
| 73 | def run(model, classifier, dataloader, args): |
| 74 | autocast = get_autocast(args.precision) |
| 75 | model = unwrap_model(model) |
| 76 | total_batch_size = dataloader.batch_size * args.world_size |
| 77 | with torch.no_grad(): |
| 78 | top1, top5, n = 0., 0., 0. |
| 79 | bar = tqdm(dataloader, unit_scale=total_batch_size) |
| 80 | for images, target in bar: |
| 81 | images = images.to(args.device) |
| 82 | target = target.to(args.device) |
| 83 | batch_size = images.size(0) |
| 84 | |
| 85 | with autocast(): |
| 86 | # predict |
| 87 | image_features = model.encode_image(images) |
| 88 | image_features = F.normalize(image_features, dim=-1) |
| 89 | logits = 100. * image_features @ classifier |
| 90 | |
| 91 | # measure accuracy |
| 92 | acc1, acc5 = accuracy(logits, target, topk=(1, 5)) |
| 93 | bar.set_description( |
| 94 | f'Acc@1 {acc1 / batch_size:.3f} Acc@5 {acc5 / batch_size:.3f}') |
| 95 | top1 += acc1 |
| 96 | top5 += acc5 |
| 97 | n += batch_size |
| 98 | del images, target, logits |
| 99 | |
| 100 | # sync top1, top5 and n |
| 101 | data = torch.tensor([top1, top5, n]).cuda() |
| 102 | dist.all_reduce(data, op=dist.ReduceOp.SUM) |
| 103 | top1, top5, n = data.tolist() |
| 104 | |
| 105 | top1 = (top1 / n) |
| 106 | top5 = (top5 / n) |
| 107 | return top1, top5 |
| 108 | |
| 109 | |
| 110 | def zero_shot_eval(model, data, epoch, args): |
no test coverage detected