test_mixins.py 7.03 KB
Newer Older
Kai Chen's avatar
Kai Chen committed
1
2
3
4
5
6
7
8
9
from mmdet.core import (bbox2roi, bbox_mapping, merge_aug_proposals,
                        merge_aug_bboxes, merge_aug_masks, multiclass_nms)


class RPNTestMixin(object):

    def simple_test_rpn(self, x, img_meta, rpn_test_cfg):
        rpn_outs = self.rpn_head(x)
        proposal_inputs = rpn_outs + (img_meta, rpn_test_cfg)
10
        proposal_list = self.rpn_head.get_bboxes(*proposal_inputs)
Kai Chen's avatar
Kai Chen committed
11
12
13
14
15
16
17
18
19
        return proposal_list

    def aug_test_rpn(self, feats, img_metas, rpn_test_cfg):
        imgs_per_gpu = len(img_metas[0])
        aug_proposals = [[] for _ in range(imgs_per_gpu)]
        for x, img_meta in zip(feats, img_metas):
            proposal_list = self.simple_test_rpn(x, img_meta, rpn_test_cfg)
            for i, proposals in enumerate(proposal_list):
                aug_proposals[i].append(proposals)
sty-yyj's avatar
sty-yyj committed
20
21
22
23
24
25
26
27
        # reorganize the order of 'img_metas' to match the dimensions
        # of 'aug_proposals'
        aug_img_metas = []
        for i in range(imgs_per_gpu):
            aug_img_meta = []
            for j in range(len(img_metas)):
                aug_img_meta.append(img_metas[j][i])
            aug_img_metas.append(aug_img_meta)
Kai Chen's avatar
Kai Chen committed
28
29
        # after merging, proposals will be rescaled to the original image size
        merged_proposals = [
sty-yyj's avatar
sty-yyj committed
30
31
            merge_aug_proposals(proposals, aug_img_meta, rpn_test_cfg)
            for proposals, aug_img_meta in zip(aug_proposals, aug_img_metas)
Kai Chen's avatar
Kai Chen committed
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
        ]
        return merged_proposals


class BBoxTestMixin(object):

    def simple_test_bboxes(self,
                           x,
                           img_meta,
                           proposals,
                           rcnn_test_cfg,
                           rescale=False):
        """Test only det bboxes without augmentation."""
        rois = bbox2roi(proposals)
        roi_feats = self.bbox_roi_extractor(
            x[:len(self.bbox_roi_extractor.featmap_strides)], rois)
myownskyW7's avatar
myownskyW7 committed
48
49
        if self.with_shared_head:
            roi_feats = self.shared_head(roi_feats)
Kai Chen's avatar
Kai Chen committed
50
51
52
53
54
55
56
57
58
59
        cls_score, bbox_pred = self.bbox_head(roi_feats)
        img_shape = img_meta[0]['img_shape']
        scale_factor = img_meta[0]['scale_factor']
        det_bboxes, det_labels = self.bbox_head.get_det_bboxes(
            rois,
            cls_score,
            bbox_pred,
            img_shape,
            scale_factor,
            rescale=rescale,
60
            cfg=rcnn_test_cfg)
Kai Chen's avatar
Kai Chen committed
61
62
        return det_bboxes, det_labels

63
    def aug_test_bboxes(self, feats, img_metas, proposal_list, rcnn_test_cfg):
Kai Chen's avatar
Kai Chen committed
64
65
66
67
68
69
70
        aug_bboxes = []
        aug_scores = []
        for x, img_meta in zip(feats, img_metas):
            # only one image in the batch
            img_shape = img_meta[0]['img_shape']
            scale_factor = img_meta[0]['scale_factor']
            flip = img_meta[0]['flip']
71
72
73
            # TODO more flexible
            proposals = bbox_mapping(proposal_list[0][:, :4], img_shape,
                                     scale_factor, flip)
Kai Chen's avatar
Kai Chen committed
74
75
76
77
            rois = bbox2roi([proposals])
            # recompute feature maps to save GPU memory
            roi_feats = self.bbox_roi_extractor(
                x[:len(self.bbox_roi_extractor.featmap_strides)], rois)
myownskyW7's avatar
myownskyW7 committed
78
79
            if self.with_shared_head:
                roi_feats = self.shared_head(roi_feats)
Kai Chen's avatar
Kai Chen committed
80
81
82
83
84
85
            cls_score, bbox_pred = self.bbox_head(roi_feats)
            bboxes, scores = self.bbox_head.get_det_bboxes(
                rois,
                cls_score,
                bbox_pred,
                img_shape,
86
                scale_factor,
Kai Chen's avatar
Kai Chen committed
87
                rescale=False,
88
                cfg=None)
Kai Chen's avatar
Kai Chen committed
89
90
91
92
            aug_bboxes.append(bboxes)
            aug_scores.append(scores)
        # after merging, bboxes will be rescaled to the original image size
        merged_bboxes, merged_scores = merge_aug_bboxes(
93
            aug_bboxes, aug_scores, img_metas, rcnn_test_cfg)
94
95
96
97
        det_bboxes, det_labels = multiclass_nms(merged_bboxes, merged_scores,
                                                rcnn_test_cfg.score_thr,
                                                rcnn_test_cfg.nms,
                                                rcnn_test_cfg.max_per_img)
Kai Chen's avatar
Kai Chen committed
98
99
100
101
102
103
104
105
106
107
108
109
        return det_bboxes, det_labels


class MaskTestMixin(object):

    def simple_test_mask(self,
                         x,
                         img_meta,
                         det_bboxes,
                         det_labels,
                         rescale=False):
        # image shape of the first image in the batch (only one)
110
        ori_shape = img_meta[0]['ori_shape']
Kai Chen's avatar
Kai Chen committed
111
112
113
114
115
116
        scale_factor = img_meta[0]['scale_factor']
        if det_bboxes.shape[0] == 0:
            segm_result = [[] for _ in range(self.mask_head.num_classes - 1)]
        else:
            # if det_bboxes is rescaled to the original image size, we need to
            # rescale it back to the testing scale to obtain RoIs.
117
118
            _bboxes = (
                det_bboxes[:, :4] * scale_factor if rescale else det_bboxes)
Kai Chen's avatar
Kai Chen committed
119
120
121
            mask_rois = bbox2roi([_bboxes])
            mask_feats = self.mask_roi_extractor(
                x[:len(self.mask_roi_extractor.featmap_strides)], mask_rois)
myownskyW7's avatar
myownskyW7 committed
122
123
            if self.with_shared_head:
                mask_feats = self.shared_head(mask_feats)
Kai Chen's avatar
Kai Chen committed
124
            mask_pred = self.mask_head(mask_feats)
125
126
127
128
129
            segm_result = self.mask_head.get_seg_masks(mask_pred, _bboxes,
                                                       det_labels,
                                                       self.test_cfg.rcnn,
                                                       ori_shape, scale_factor,
                                                       rescale)
Kai Chen's avatar
Kai Chen committed
130
131
        return segm_result

132
    def aug_test_mask(self, feats, img_metas, det_bboxes, det_labels):
Kai Chen's avatar
Kai Chen committed
133
134
135
136
137
138
139
140
141
142
143
144
145
146
        if det_bboxes.shape[0] == 0:
            segm_result = [[] for _ in range(self.mask_head.num_classes - 1)]
        else:
            aug_masks = []
            for x, img_meta in zip(feats, img_metas):
                img_shape = img_meta[0]['img_shape']
                scale_factor = img_meta[0]['scale_factor']
                flip = img_meta[0]['flip']
                _bboxes = bbox_mapping(det_bboxes[:, :4], img_shape,
                                       scale_factor, flip)
                mask_rois = bbox2roi([_bboxes])
                mask_feats = self.mask_roi_extractor(
                    x[:len(self.mask_roi_extractor.featmap_strides)],
                    mask_rois)
myownskyW7's avatar
myownskyW7 committed
147
148
                if self.with_shared_head:
                    mask_feats = self.shared_head(mask_feats)
Kai Chen's avatar
Kai Chen committed
149
150
151
152
                mask_pred = self.mask_head(mask_feats)
                # convert to numpy array to save memory
                aug_masks.append(mask_pred.sigmoid().cpu().numpy())
            merged_masks = merge_aug_masks(aug_masks, img_metas,
153
154
155
                                           self.test_cfg.rcnn)

            ori_shape = img_metas[0][0]['ori_shape']
Kai Chen's avatar
Kai Chen committed
156
            segm_result = self.mask_head.get_seg_masks(
pangjm's avatar
pangjm committed
157
158
159
160
161
162
163
                merged_masks,
                det_bboxes,
                det_labels,
                self.test_cfg.rcnn,
                ori_shape,
                scale_factor=1.0,
                rescale=False)
Kai Chen's avatar
Kai Chen committed
164
        return segm_result