(cfg, model, is_last=False, use_tta=False)
| 195 | |
| 196 | |
| 197 | def do_test(cfg, model, is_last=False, use_tta=False): |
| 198 | if not cfg.TEST.ENABLED: |
| 199 | LOG.warning("Test is disabled.") |
| 200 | return {} |
| 201 | |
| 202 | dataset_names = [cfg.DATASETS.TEST.NAME] # NOTE: only support single test dataset for now. |
| 203 | |
| 204 | if use_tta: |
| 205 | LOG.info("Starting inference with test-time augmentation.") |
| 206 | if isinstance(model, DistributedDataParallel): |
| 207 | model.module.postprocess_in_inference = False |
| 208 | else: |
| 209 | model.postprocess_in_inference = False |
| 210 | model = build_tta_model(cfg, model) |
| 211 | |
| 212 | test_results = OrderedDict() |
| 213 | for dataset_name in dataset_names: |
| 214 | # output directory for this dataset. |
| 215 | dset_output_dir = get_inference_output_dir(dataset_name, is_last=is_last, use_tta=use_tta) |
| 216 | |
| 217 | # What evaluators are used for this dataset? |
| 218 | evaluator_names = MetadataCatalog.get(dataset_name).evaluators |
| 219 | evaluators = [] |
| 220 | for evaluator_name in evaluator_names: |
| 221 | evaluator = get_evaluator(cfg, dataset_name, evaluator_name, dset_output_dir) |
| 222 | evaluators.append(evaluator) |
| 223 | evaluator = DatasetEvaluators(evaluators) |
| 224 | |
| 225 | mapper = get_dataset_mapper(cfg, is_train=False) |
| 226 | dataloader, dataset_dicts = build_test_dataloader(cfg, dataset_name, mapper) |
| 227 | |
| 228 | per_dataset_results = inference_on_dataset(model, dataloader, evaluator) |
| 229 | if use_tta: |
| 230 | per_dataset_results = OrderedDict({k + '-tta': v for k, v in per_dataset_results.items()}) |
| 231 | test_results[dataset_name] = per_dataset_results |
| 232 | |
| 233 | if cfg.VIS.PREDICTIONS_ENABLED and d2_comm.is_main_process(): |
| 234 | visualizer_names = MetadataCatalog.get(dataset_name).pred_visualizers |
| 235 | # Randomly (but deterministically) select what samples to visualize. |
| 236 | # The samples are shared across all visualizers and iterations. |
| 237 | sampled_dataset_dicts, inds = random_sample_dataset_dicts( |
| 238 | dataset_name, num_samples=cfg.VIS.PREDICTIONS_MAX_NUM_SAMPLES |
| 239 | ) |
| 240 | |
| 241 | viz_images = defaultdict(dict) |
| 242 | for viz_name in visualizer_names: |
| 243 | LOG.info(f"Running prediction visualizer: {viz_name}") |
| 244 | visualizer = get_predictions_visualizer(cfg, viz_name, dataset_name, dset_output_dir) |
| 245 | for x in tqdm(sampled_dataset_dicts): |
| 246 | sample_id = x['sample_id'] |
| 247 | viz_images[sample_id].update(visualizer.visualize(x)) |
| 248 | |
| 249 | save_vis(viz_images, dset_output_dir, "visualization") |
| 250 | |
| 251 | if cfg.WANDB.ENABLED: |
| 252 | LOG.info(f"Uploading prediction visualization to W&B: {dataset_name}") |
| 253 | for sample_id in viz_images.keys(): |
| 254 | viz_images[sample_id] = mosaic(list(viz_images[sample_id].values())) |
no test coverage detected