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
|