| 39 | |
| 40 | |
| 41 | class COCOCenterNetDatasetMapper: |
| 42 | |
| 43 | def __init__(self, cfg, is_train=True): |
| 44 | if cfg.INPUT.CROP.ENABLED and is_train: |
| 45 | self.crop_gen = T.RandomCrop(cfg.INPUT.CROP.TYPE, cfg.INPUT.CROP.SIZE) |
| 46 | logging.getLogger(__name__).info("CropGen used in training: " + str(self.crop_gen)) |
| 47 | else: |
| 48 | self.crop_gen = None |
| 49 | |
| 50 | self.tfm_gens = build_COCO_transform_gen(cfg, is_train) |
| 51 | |
| 52 | # fmt: off |
| 53 | self.img_format = cfg.INPUT.FORMAT |
| 54 | self.mask_on = cfg.MODEL.MASK_ON |
| 55 | self.mask_format = cfg.INPUT.MASK_FORMAT |
| 56 | self.keypoint_on = cfg.MODEL.KEYPOINT_ON |
| 57 | self.down_ratio = cfg.MODEL.CENTERNET.DOWN_RATIO |
| 58 | self.num_classes = cfg.MODEL.CENTERNET.NUM_CLASSES |
| 59 | self.load_proposals = cfg.MODEL.LOAD_PROPOSALS |
| 60 | # fmt: on |
| 61 | if self.keypoint_on and is_train: |
| 62 | # Flip only makes sense in training |
| 63 | self.keypoint_hflip_indices = utils.create_keypoint_hflip_indices(cfg.DATASETS.TRAIN) |
| 64 | else: |
| 65 | self.keypoint_hflip_indices = None |
| 66 | |
| 67 | if self.load_proposals: |
| 68 | self.min_box_side_len = cfg.MODEL.PROPOSAL_GENERATOR.MIN_SIZE |
| 69 | self.proposal_topk = ( |
| 70 | cfg.DATASETS.PRECOMPUTED_PROPOSAL_TOPK_TRAIN |
| 71 | if is_train |
| 72 | else cfg.DATASETS.PRECOMPUTED_PROPOSAL_TOPK_TEST |
| 73 | ) |
| 74 | self.is_train = is_train |
| 75 | |
| 76 | def __call__(self, dataset_dict): |
| 77 | dataset_dict = copy.deepcopy(dataset_dict) |
| 78 | image = utils.read_image(dataset_dict["file_name"], format=self.img_format) |
| 79 | utils.check_image_size(dataset_dict, image) |
| 80 | |
| 81 | image, transforms = T.apply_transform_gens(self.tfm_gens, image) |
| 82 | |
| 83 | image_shape = image.shape[:2] # h, w |
| 84 | dataset_dict["image"] = torch.as_tensor(np.ascontiguousarray(image.transpose(2, 0, 1))) |
| 85 | |
| 86 | if not self.is_train: |
| 87 | # USER: Modify this if you want to keep them for some reason. |
| 88 | dataset_dict.pop("annotations", None) |
| 89 | dataset_dict.pop("sem_seg_file_name", None) |
| 90 | return dataset_dict |
| 91 | |
| 92 | if "annotations" in dataset_dict: |
| 93 | # USER: Modify this if you want to keep them for some reason. |
| 94 | for anno in dataset_dict["annotations"]: |
| 95 | if not (self.mask_on or self.keypoint_on): |
| 96 | anno.pop("segmentation", None) |
| 97 | if not self.keypoint_on: |
| 98 | anno.pop("keypoints", None) |
no outgoing calls
no test coverage detected