MCPcopy Create free account
hub / github.com/microsoft/Cream / evaluate

Function evaluate

AutoFormer/supernet_engine.py:115–160  ·  view source on GitHub ↗
(data_loader, model, device, amp=True, choices=None, mode='super', retrain_config=None)

Source from the content-addressed store, hash-verified

113
114@torch.no_grad()
115def 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()}

Callers 2

mainFunction · 0.90
is_legalMethod · 0.90

Calls 10

log_everyMethod · 0.95
updateMethod · 0.95
accuracyFunction · 0.90
sample_configsFunction · 0.85
formatMethod · 0.80
toMethod · 0.80
printFunction · 0.50
set_sample_configMethod · 0.45

Tested by

no test coverage detected