File size: 17,603 Bytes
6ca1e94
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
import torch
import torch.nn as nn
import torchvision.ops as ops
from torchvision.models import resnet50, ResNet50_Weights

class DetectionHead(nn.Module):
    def __init__(self, num_classes=20, pretrained=True):
        super().__init__()

        self.num_classes = num_classes

        weights = ResNet50_Weights.DEFAULT if pretrained else None
        self.conv5_x = resnet50(weights=weights).layer4

        self.fc_cls = nn.Linear(2048, num_classes + 1)
        self.fc_reg = nn.Linear(2048, num_classes * 4)

        self.fc_cls.weight.data.normal_(0, 0.01)
        self.fc_reg.weight.data.normal_(0, 0.001)
        self.fc_cls.bias.data.zero_()
        self.fc_reg.bias.data.zero_()

        for module in self.conv5_x.modules():
            if isinstance(module, nn.BatchNorm2d):
                module.weight.requires_grad = False
                module.bias.requires_grad = False

    def train(self, mode=True):
            super().train(mode)
    
            # requires_grad=False only freezes the affine params
            # training mode may again put the batchnorm layers in training mode, which is not what we want
            for module in self.conv5_x.modules():
                if isinstance(module, nn.BatchNorm2d):
                    module.eval()
    
            return self

    def forward(self, pooled_batch):
        # pooled_batch: list of B tensors [N_i, 1024, 7, 7] from RoIPool
        counts = [pooled.shape[0] for pooled in pooled_batch]

        # No proposals survived anywhere in the batch -- torch.cat would still work, but
        # conv5_x on a 0-length batch is pointless, so return correctly shaped empties
        if sum(counts) == 0:
            cls_logits = [pooled.new_zeros((0, self.num_classes + 1)) for pooled in pooled_batch]
            box_deltas = [pooled.new_zeros((0, self.num_classes * 4)) for pooled in pooled_batch]
            return cls_logits, box_deltas

        # One batched pass over every RoI in the batch, then split back per image
        x = torch.cat(pooled_batch, dim=0)           # [sumN, 1024, 7, 7]
        x = self.conv5_x(x)                          # [sumN, 2048, 4, 4]
        x = x.mean(dim=(2, 3))                       # global average pool -> [sumN, 2048]

        cls_logits = self.fc_cls(x)                  # [sumN, 21] raw logits
        box_deltas = self.fc_reg(x)                  # [sumN, 80] class-specific deltas

        return list(cls_logits.split(counts)), list(box_deltas.split(counts))


class DetectionLoss(nn.Module):
    def __init__(self):
        super().__init__()

    def forward(self, feature_maps, batch_proposals, batch_gt_boxes, batch_gt_labels,
                batch_img_height, batch_img_width, roi_pool, detection_head,
                num_samples=64, pos_fraction=0.25, background_label=20):
        # feature_maps:   [B, 1024, H_f, W_f] from Backbone
        # batch_proposals: list of B x [N_i, 4] CORNER absolute px, from RegionProposalNetwork
        # batch_gt_boxes:  list of B x [M_i, 4] CORNER absolute px
        # batch_gt_labels: list of B x [M_i] in 0..19
        device = feature_maps.device
        batch_size = len(batch_proposals)

        batch_sampled_proposals = []
        batch_labels = []
        batch_reg_targets = []

        # 64 sampled RoIs per image instead of all ~2000 proposals.
        for i in range(batch_size):
            # Step 2 trains the detector on the step-1 RPN's proposals as FIXED input.
            # reg_targets are built from these boxes, so without the detach the regression
            # target itself would be differentiable and smooth_l1_loss would backpropagate
            # into the RPN.
            proposals = batch_proposals[i].detach()
            gt_boxes = batch_gt_boxes[i].to(device)
            gt_labels = batch_gt_labels[i].to(device)

            labels, matched_gt_idx = self.assign_proposal_labels(proposals, gt_boxes, gt_labels, background_label=background_label)
            sampled_idx = self.sample_proposals(labels, num_samples, pos_fraction, background_label)

            sampled_proposals = proposals[sampled_idx]
            sampled_labels = labels[sampled_idx]

            if gt_boxes.shape[0] == 0:
                # Every label is background here, so these targets are never read.
                # Indexing gt_boxes with matched_gt_idx would be out of bounds.
                reg_targets = torch.zeros((sampled_idx.numel(), 4), dtype=sampled_proposals.dtype, device=device)   # unused placeholder gets discarded by selecting only foreground labels
            else:
                matched_gt_boxes = gt_boxes[matched_gt_idx[sampled_idx]]
                reg_targets = self.encode_box_targets(sampled_proposals, matched_gt_boxes)

            batch_sampled_proposals.append(sampled_proposals)
            batch_labels.append(sampled_labels)
            batch_reg_targets.append(reg_targets)

        pooled_batch = roi_pool(feature_maps, batch_sampled_proposals, batch_img_height, batch_img_width)
        batch_cls_logits, batch_box_deltas = detection_head(pooled_batch)

        total_cls_loss = 0.0
        total_reg_loss = 0.0

        for i in range(batch_size):
            cls_logits = batch_cls_logits[i]
            box_deltas = batch_box_deltas[i]
            labels = batch_labels[i]
            reg_targets = batch_reg_targets[i]

            num_sampled = labels.numel()
            if num_sampled == 0:
                # Nothing survived sampling for this image; it contributes zero to both terms
                continue

            cls_loss = nn.functional.cross_entropy(cls_logits, labels, reduction='mean')

            # For reg_loss, selecting sampled proposals that are not background (i.e., positive samples) is necessary because the regression loss is only computed for foreground classes.
            positive_mask = labels != background_label

            if positive_mask.sum() == 0:
                # smooth_l1_loss on an empty tensor returns nan and would poison the batch
                reg_loss = torch.tensor(0.0, dtype=box_deltas.dtype, device=device)
            else:
                # box_deltas is [N, 4 * 20] 
                num_classes = box_deltas.shape[1] // 4
                positive_labels = labels[positive_mask]
                positive_deltas = box_deltas[positive_mask].view(-1, num_classes, 4)
                # Select the predicted deltas for each proposal based on proposal idx and proposal_label(for that prop_idx)
                # torch.arange(positive_labels.numel(), device=device) creates [0, ..., P = num_positive_proposal], positive_labels = [cls_label_prop_1, cls_label_prop_2, ..., cls_label_prop_P]
                predicted_deltas = positive_deltas[torch.arange(positive_labels.numel(), device=device), positive_labels]

                reg_loss = nn.functional.smooth_l1_loss(predicted_deltas, reg_targets[positive_mask], reduction='sum')

                reg_loss = reg_loss / num_sampled

            total_cls_loss += cls_loss
            total_reg_loss += reg_loss

        return total_cls_loss / batch_size, total_reg_loss / batch_size

    def compute_iou_matrix_corners(self, boxes_a, boxes_b):
        a_x1, a_y1, a_x2, a_y2 = boxes_a[:, 0], boxes_a[:, 1], boxes_a[:, 2], boxes_a[:, 3]
        b_x1, b_y1, b_x2, b_y2 = boxes_b[:, 0], boxes_b[:, 1], boxes_b[:, 2], boxes_b[:, 3]

        inter_x1 = torch.max(a_x1.unsqueeze(1), b_x1.unsqueeze(0))
        inter_y1 = torch.max(a_y1.unsqueeze(1), b_y1.unsqueeze(0))
        inter_x2 = torch.min(a_x2.unsqueeze(1), b_x2.unsqueeze(0))
        inter_y2 = torch.min(a_y2.unsqueeze(1), b_y2.unsqueeze(0))

        inter_area = torch.clamp(inter_x2 - inter_x1, min=0) * torch.clamp(inter_y2 - inter_y1, min=0)

        area_a = (a_x2 - a_x1) * (a_y2 - a_y1)
        area_b = (b_x2 - b_x1) * (b_y2 - b_y1)

        union_area = area_a.unsqueeze(1) + area_b.unsqueeze(0) - inter_area

        return inter_area / union_area

    def assign_proposal_labels(self, proposals, gt_boxes, gt_labels, pos_iou_thresh=0.5, neg_iou_lo=0.1, background_label=20):
        # proposals: [N, 4] CORNER absolute px, straight from RegionProposalNetwork
        # gt_boxes:  [M, 4] CORNER absolute px    gt_labels: [M] in 0..19
        # Returns labels [N] (0..19 foreground, 20 background, -1 ignore) and
        # matched_gt_idx [N] (index of the best-overlapping gt box for each proposal).
        num_proposals = proposals.shape[0]

        if gt_boxes.shape[0] == 0:
            # No objects in this image, so every proposal is background
            labels = torch.full((num_proposals,), background_label, dtype=torch.long, device=proposals.device)
            matched_gt_idx = torch.zeros((num_proposals,), dtype=torch.long, device=proposals.device)  # unused placeholder
            return labels, matched_gt_idx

        iou_matrix = self.compute_iou_matrix_corners(proposals, gt_boxes)
        max_iou_per_proposal, matched_gt_idx = iou_matrix.max(dim=1)

        labels = torch.full((num_proposals,), -1, dtype=torch.long, device=proposals.device)
        labels[max_iou_per_proposal >= neg_iou_lo] = background_label

        positive_mask = max_iou_per_proposal >= pos_iou_thresh
        labels[positive_mask] = gt_labels[matched_gt_idx[positive_mask]]

        return labels, matched_gt_idx

    def sample_proposals(self, labels, num_samples=64, pos_fraction=0.25, background_label=20):
        # 64 per image, 25% of them positive.
        positive_idx = torch.where((labels >= 0) & (labels != background_label))[0]
        negative_idx = torch.where(labels == background_label)[0]

        num_pos = min(int(num_samples * pos_fraction), positive_idx.numel())
        num_neg = min(num_samples - num_pos, negative_idx.numel())

        perm_pos = torch.randperm(positive_idx.numel(), device=labels.device)[:num_pos]
        perm_neg = torch.randperm(negative_idx.numel(), device=labels.device)[:num_neg]

        sampled_pos_idx = positive_idx[perm_pos]
        sampled_neg_idx = negative_idx[perm_neg]

        # positives first, then negatives -- ignored proposals never enter the result
        return torch.cat([sampled_pos_idx, sampled_neg_idx])

    def corners_to_center(self, boxes):
        # (x1, y1, x2, y2) -> (x_c, y_c, w, h)
        x1, y1, x2, y2 = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]
        w = x2 - x1
        h = y2 - y1
        x_c = x1 + w / 2
        y_c = y1 + h / 2
        return torch.stack([x_c, y_c, w, h], dim=1)

    def encode_box_targets(self, proposals, gt_boxes, delta_std=(0.1, 0.1, 0.2, 0.2)):
        proposals_centered = self.corners_to_center(proposals)
        gt_centered = self.corners_to_center(gt_boxes)

        p_xc, p_yc, p_w, p_h = proposals_centered[:, 0], proposals_centered[:, 1], proposals_centered[:, 2], proposals_centered[:, 3]
        gt_xc, gt_yc, gt_w, gt_h = gt_centered[:, 0], gt_centered[:, 1], gt_centered[:, 2], gt_centered[:, 3]

        target_dx = (gt_xc - p_xc) / p_w
        target_dy = (gt_yc - p_yc) / p_h
        target_dw = torch.log(gt_w / p_w)
        target_dh = torch.log(gt_h / p_h)

        target_deltas = torch.stack((target_dx, target_dy, target_dw, target_dh), dim=1)

        # Standard Deviations of the regression targets
        std = torch.tensor(delta_std, dtype=target_deltas.dtype, device=target_deltas.device)

        # Based on the Dataset, the mean of target_deltas is appx 0 and std dev (0.1, 0.1, 0.2, 0.2)
        # So, by (target_deltas - 0)/std we are normalizing the regression targets to have zero mean and unit variance, as described in the Fast R-CNN paper.
        return target_deltas / std


class DetectionNet(nn.Module):
    def __init__(self, detection_head, background_label=20, delta_std=(0.1, 0.1, 0.2, 0.2),
                 score_thresh=0.05, nms_iou_thresh=0.5, max_detections_per_image=100,
                 min_box_size=1.0):
        super().__init__()

        self.detection_head = detection_head
        self.background_label = background_label
        self.delta_std = delta_std

        # Inference-time detection settings. score_thresh=0.05 and 100 detections per
        # image are the Fast R-CNN evaluation conventions; both are deliberately low
        # bars, since AP wants the full precision/recall curve rather than one
        # conservative operating point.
        self.score_thresh = score_thresh
        self.nms_iou_thresh = nms_iou_thresh
        self.max_detections_per_image = max_detections_per_image
        self.min_box_size = min_box_size

    def corners_to_center(self, boxes):
        # (x1, y1, x2, y2) -> (x_c, y_c, w, h)
        x1, y1, x2, y2 = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]
        w = x2 - x1
        h = y2 - y1
        x_c = x1 + w / 2
        y_c = y1 + h / 2
        return torch.stack([x_c, y_c, w, h], dim=1)

    def select_foreground(self, cls_logits):
        # score per proposal per class, excluding the background class (last column)
        softmax_scores = nn.functional.softmax(cls_logits, dim=-1)      # [N, num_classes + 1]
        fg_scores = softmax_scores[:, :self.background_label]           # [N, num_classes] -- background class is not a foreground class

        proposal_idx, labels = (fg_scores > self.score_thresh).nonzero(as_tuple=True)       # proposal_idx: [K], labels: [K] -- K is the number of proposals that survived score_thresh for at least one class
        scores = fg_scores[proposal_idx, labels]

        return proposal_idx, labels, scores

    def decode_box_deltas(self, proposals_center, box_deltas):
        # proposals_center: [N, 4] (x_c, y_c, w, h)
        # box_deltas: [N, 4] (dx, dy, dw, dh), normalized by self.delta_std (Fast R-CNN convention, matches DetectionLoss.encode_box_targets) -- un-normalize before applying the inverse transform
        std = torch.tensor(self.delta_std, dtype=box_deltas.dtype, device=box_deltas.device)
        dx, dy, dw, dh = (box_deltas * std).unbind(dim=1)

        xc_p, yc_p, w_p, h_p = proposals_center[:, 0], proposals_center[:, 1], proposals_center[:, 2], proposals_center[:, 3]

        xc = dx * w_p + xc_p
        yc = dy * h_p + yc_p
        w = torch.exp(dw) * w_p
        h = torch.exp(dh) * h_p

        xmin = xc - w / 2
        ymin = yc - h / 2
        xmax = xc + w / 2
        ymax = yc + h / 2

        decoded_boxes = torch.stack((xmin, ymin, xmax, ymax), dim=1)
        return decoded_boxes

    def clip_boxes_to_image(self, decoded_boxes, img_height, img_width):
        # decoded_boxes: [N, 4] CORNER absolute px -- same clamp as RegionProposalNetwork.clip_boxes_to_image
        xmin = torch.clamp(decoded_boxes[:, 0], min=0, max=img_width - 1)
        ymin = torch.clamp(decoded_boxes[:, 1], min=0, max=img_height - 1)
        xmax = torch.clamp(decoded_boxes[:, 2], min=0, max=img_width - 1)
        ymax = torch.clamp(decoded_boxes[:, 3], min=0, max=img_height - 1)

        clipped_boxes = torch.stack((xmin, ymin, xmax, ymax), dim=1)
        return clipped_boxes

    def forward(self, rpn_proposals, pooled_proposals, img_sizes_before_pad):
        # rpn_proposals:    list of B x [N_i, 4] CORNER absolute px, from RegionProposalNetwork
        # pooled_proposals: list of B x [N_i, 1024, 7, 7], RoIPool output for the same proposals
        batch_cls_logits, batch_box_deltas = self.detection_head(pooled_proposals)

        batch_size = len(rpn_proposals)

        labels_list = []
        scores_list = []
        boxes_list = []

        for i in range(batch_size):
            img_height, img_width = img_sizes_before_pad[i]

            cls_logits = batch_cls_logits[i]
            box_deltas = batch_box_deltas[i]
            proposals = rpn_proposals[i]

            proposal_idx, labels, scores = self.select_foreground(cls_logits)

            if proposal_idx.numel() == 0:
                # Nothing in this image cleared score_thresh for any class
                labels_list.append(labels)
                scores_list.append(scores)
                boxes_list.append(proposals.new_zeros((0, 4)))
                continue

            # box_deltas is [N, 4 * num_classes] -- reshape to [N, num_classes, 4] 
            num_classes = box_deltas.shape[1] // 4
            selected_proposals = proposals[proposal_idx]
            predicted_deltas = box_deltas.view(-1, num_classes, 4)[proposal_idx, labels]

            proposals_center = self.corners_to_center(selected_proposals)
            decoded_boxes = self.decode_box_deltas(proposals_center, predicted_deltas)
            clipped_boxes = self.clip_boxes_to_image(decoded_boxes, img_height, img_width)

            # Remove boxes that are too small to be valid detections
            widths = clipped_boxes[:, 2] - clipped_boxes[:, 0]
            heights = clipped_boxes[:, 3] - clipped_boxes[:, 1]
            keep = (widths >= self.min_box_size) & (heights >= self.min_box_size)
            clipped_boxes, labels, scores = clipped_boxes[keep], labels[keep], scores[keep]


            if labels.numel() == 0:
                labels_list.append(labels)
                scores_list.append(scores)
                boxes_list.append(clipped_boxes)
                continue

            # Per-CLASS NMS -- boxes of different classes must not suppress each other
            keep_idx = ops.batched_nms(clipped_boxes, scores, labels, self.nms_iou_thresh)
            keep_idx = keep_idx[:self.max_detections_per_image]

            labels_list.append(labels[keep_idx])
            scores_list.append(scores[keep_idx])
            boxes_list.append(clipped_boxes[keep_idx])

        return labels_list, scores_list, boxes_list