(cfg)
| 22 | |
| 23 | @hydra.main(config_path="../configs/", config_name="visualize_dataloader") |
| 24 | def main(cfg): |
| 25 | setup(cfg) |
| 26 | dataset_names = register_datasets(cfg) |
| 27 | if cfg.ONLY_REGISTER_DATASETS: |
| 28 | return {}, cfg |
| 29 | LOG.info(f"Registered {len(dataset_names)} datasets:" + '\n\t' + '\n\t'.join(dataset_names)) |
| 30 | |
| 31 | if cfg.USE_TEST: |
| 32 | dataset_name = cfg.DATASETS.TEST.NAME |
| 33 | mapper = get_dataset_mapper(cfg, is_train=False) |
| 34 | dataloader, _ = build_test_dataloader(cfg, dataset_name, mapper=mapper) |
| 35 | else: |
| 36 | mapper = get_dataset_mapper(cfg, is_train=True) |
| 37 | dataloader, _ = build_train_dataloader(cfg, mapper=mapper) |
| 38 | |
| 39 | visualizer_names = MetadataCatalog.get(cfg.DATASETS.TRAIN.NAME).loader_visualizers |
| 40 | for batch_idx, batch in tqdm(enumerate(dataloader)): |
| 41 | viz_images = defaultdict(dict) |
| 42 | LOG.info("Press any key to continue, press 'q' to quit.") |
| 43 | for viz_name in visualizer_names: |
| 44 | viz = get_dataloader_visualizer(cfg, viz_name, cfg.DATASETS.TRAIN.NAME) |
| 45 | for idx, x in enumerate(batch): |
| 46 | viz_images[idx].update(viz.visualize(x)) |
| 47 | |
| 48 | for k in range(len(batch)): |
| 49 | gt_viz = mosaic(list(viz_images[k].values())) |
| 50 | cv2.imshow("dataloader", gt_viz[:, :, ::-1]) |
| 51 | |
| 52 | if cv2.waitKey(0) & 0xFF == ord('q'): |
| 53 | sys.exit() |
| 54 | |
| 55 | |
| 56 | if __name__ == '__main__': |
no test coverage detected