(data_loader, model, device, amp=True, choices=None, mode='super', retrain_config=None)
| 113 | |
| 114 | @torch.no_grad() |
| 115 | def evaluate(data_loader, model, device, amp=True, choices=None, mode='super', retrain_config=None): |
| 116 | criterion = torch.nn.CrossEntropyLoss() |
| 117 | |
| 118 | metric_logger = utils.MetricLogger(delimiter=" ") |
| 119 | header = 'Test:' |
| 120 | |
| 121 | # switch to evaluation mode |
| 122 | model.eval() |
| 123 | if mode == 'super': |
| 124 | config = sample_configs(choices=choices) |
| 125 | model_module = unwrap_model(model) |
| 126 | model_module.set_sample_config(config=config) |
| 127 | else: |
| 128 | config = retrain_config |
| 129 | model_module = unwrap_model(model) |
| 130 | model_module.set_sample_config(config=config) |
| 131 | |
| 132 | |
| 133 | print("sampled model config: {}".format(config)) |
| 134 | parameters = model_module.get_sampled_params_numel(config) |
| 135 | print("sampled model parameters: {}".format(parameters)) |
| 136 | |
| 137 | for images, target in metric_logger.log_every(data_loader, 10, header): |
| 138 | images = images.to(device, non_blocking=True) |
| 139 | target = target.to(device, non_blocking=True) |
| 140 | # compute output |
| 141 | if amp: |
| 142 | with torch.cuda.amp.autocast(): |
| 143 | output = model(images) |
| 144 | loss = criterion(output, target) |
| 145 | else: |
| 146 | output = model(images) |
| 147 | loss = criterion(output, target) |
| 148 | |
| 149 | acc1, acc5 = accuracy(output, target, topk=(1, 5)) |
| 150 | |
| 151 | batch_size = images.shape[0] |
| 152 | metric_logger.update(loss=loss.item()) |
| 153 | metric_logger.meters['acc1'].update(acc1.item(), n=batch_size) |
| 154 | metric_logger.meters['acc5'].update(acc5.item(), n=batch_size) |
| 155 | # gather the stats from all processes |
| 156 | metric_logger.synchronize_between_processes() |
| 157 | print('* Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f} loss {losses.global_avg:.3f}' |
| 158 | .format(top1=metric_logger.acc1, top5=metric_logger.acc5, losses=metric_logger.loss)) |
| 159 | |
| 160 | return {k: meter.global_avg for k, meter in metric_logger.meters.items()} |
no test coverage detected