MCPcopy Create free account
hub / github.com/HYUNJS/SGT / COCOCenterNetDatasetMapper

Class COCOCenterNetDatasetMapper

projects/CenterNet/centernet/dataset_mapper.py:41–112  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

39
40
41class 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)

Callers 2

build_train_loaderMethod · 0.90
build_test_loaderMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected