(model,
criterion,
postprocessors,
data_loader,
device,
output_dir,
wo_class_error=False,
tmpdir=None,
gpu_collect=False,
args=None,
logger=None)
| 181 | |
| 182 | @torch.no_grad() |
| 183 | def evaluate(model, |
| 184 | criterion, |
| 185 | postprocessors, |
| 186 | data_loader, |
| 187 | device, |
| 188 | output_dir, |
| 189 | wo_class_error=False, |
| 190 | tmpdir=None, |
| 191 | gpu_collect=False, |
| 192 | args=None, |
| 193 | logger=None): |
| 194 | try: |
| 195 | need_tgt_for_training = args.use_dn |
| 196 | except: |
| 197 | need_tgt_for_training = False |
| 198 | model.eval() |
| 199 | criterion.eval() |
| 200 | |
| 201 | metric_logger = utils.MetricLogger(delimiter=' ') |
| 202 | if not wo_class_error: |
| 203 | metric_logger.add_meter( |
| 204 | 'class_error', utils.SmoothedValue(window_size=1, |
| 205 | fmt='{value:.2f}')) |
| 206 | header = 'Test:' |
| 207 | iou_types = tuple(k for k in ('bbox', 'keypoints')) |
| 208 | try: |
| 209 | useCats = args.useCats |
| 210 | except: |
| 211 | useCats = True |
| 212 | if not useCats: |
| 213 | print('useCats: {} !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!'.format( |
| 214 | useCats)) |
| 215 | |
| 216 | _cnt = 0 |
| 217 | results = [] |
| 218 | dataset = data_loader.dataset |
| 219 | rank, world_size = get_dist_info() |
| 220 | |
| 221 | if rank == 0: |
| 222 | # Check if tmpdir is valid for cpu_collect |
| 223 | if (not gpu_collect) and (tmpdir is not None and osp.exists(tmpdir)): |
| 224 | raise OSError((f'The tmpdir {tmpdir} already exists.', |
| 225 | ' Since tmpdir will be deleted after testing,', |
| 226 | ' please make sure you specify an empty one.')) |
| 227 | prog_bar = mmcv.ProgressBar(len(dataset)) |
| 228 | time.sleep(2) |
| 229 | # i=0 |
| 230 | cur_sample_idx = 0 |
| 231 | eval_result = {} |
| 232 | # print() |
| 233 | cur_eval_result_list = [] |
| 234 | rank, world_size = get_dist_info() |
| 235 | |
| 236 | for data_batch in metric_logger.log_every( |
| 237 | data_loader, 10, header, logger=logger): |
| 238 | # i = i+1 |
| 239 | with torch.cuda.amp.autocast(enabled=args.amp): |
| 240 | if need_tgt_for_training: |
no test coverage detected