MCPcopy Create free account
hub / github.com/TRI-ML/dd3d / do_test

Function do_test

scripts/train.py:197–274  ·  view source on GitHub ↗
(cfg, model, is_last=False, use_tta=False)

Source from the content-addressed store, hash-verified

195
196
197def 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()))

Callers 2

mainFunction · 0.85
do_trainFunction · 0.85

Calls 13

build_tta_modelFunction · 0.90
get_inference_output_dirFunction · 0.90
get_evaluatorFunction · 0.90
get_dataset_mapperFunction · 0.90
build_test_dataloaderFunction · 0.90
save_visFunction · 0.90
mosaicFunction · 0.90
flatten_dictFunction · 0.90
log_nested_dictFunction · 0.90
print_test_resultsFunction · 0.90

Tested by

no test coverage detected