Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import logging4from typing import List, Optional, Tuple5import torch6from fvcore.nn import sigmoid_focal_loss_jit7from torch import nn8from torch.nn import functional as F9 10from detectron2.layers import ShapeSpec, batched_nms11from detectron2.structures import Boxes, ImageList, Instances, pairwise_point_box_distance12from detectron2.utils.events import get_event_storage13 14from ..anchor_generator import DefaultAnchorGenerator15from ..backbone import Backbone16from ..box_regression import Box2BoxTransformLinear, _dense_box_regression_loss17from .dense_detector import DenseDetector18from .retinanet import RetinaNetHead19 20__all__ = ["FCOS"]21 22logger = logging.getLogger(__name__)23 24 25class FCOS(DenseDetector):26 """27 Implement FCOS in :paper:`fcos`.28 """29 30 def __init__(31 self,32 *,33 backbone: Backbone,34 head: nn.Module,35 head_in_features: Optional[List[str]] = None,36 box2box_transform=None,37 num_classes,38 center_sampling_radius: float = 1.5,39 focal_loss_alpha=0.25,40 focal_loss_gamma=2.0,41 test_score_thresh=0.2,42 test_topk_candidates=1000,43 test_nms_thresh=0.6,44 max_detections_per_image=100,45 pixel_mean,46 pixel_std,47 ):48 """49 Args:50 center_sampling_radius: radius of the "center" of a groundtruth box,51 within which all anchor points are labeled positive.52 Other arguments mean the same as in :class:`RetinaNet`.53 """54 super().__init__(55 backbone, head, head_in_features, pixel_mean=pixel_mean, pixel_std=pixel_std56 )57 58 self.num_classes = num_classes59 60 # FCOS uses one anchor point per location.61 # We represent the anchor point by a box whose size equals the anchor stride.62 feature_shapes = backbone.output_shape()63 fpn_strides = [feature_shapes[k].stride for k in self.head_in_features]64 self.anchor_generator = DefaultAnchorGenerator(65 sizes=[[k] for k in fpn_strides], aspect_ratios=[1.0], strides=fpn_strides66 )67 68 # FCOS parameterizes box regression by a linear transform,69 # where predictions are normalized by anchor stride (equal to anchor size).70 if box2box_transform is None:71 box2box_transform = Box2BoxTransformLinear(normalize_by_size=True)72 self.box2box_transform = box2box_transform73 74 self.center_sampling_radius = float(center_sampling_radius)75 76 # Loss parameters:77 self.focal_loss_alpha = focal_loss_alpha78 self.focal_loss_gamma = focal_loss_gamma79 80 # Inference parameters:81 self.test_score_thresh = test_score_thresh82 self.test_topk_candidates = test_topk_candidates83 self.test_nms_thresh = test_nms_thresh84 self.max_detections_per_image = max_detections_per_image85 86 def forward_training(self, images, features, predictions, gt_instances):87 # Transpose the Hi*Wi*A dimension to the middle:88 pred_logits, pred_anchor_deltas, pred_centerness = self._transpose_dense_predictions(89 predictions, [self.num_classes, 4, 1]90 )91 anchors = self.anchor_generator(features)92 gt_labels, gt_boxes = self.label_anchors(anchors, gt_instances)93 return self.losses(94 anchors, pred_logits, gt_labels, pred_anchor_deltas, gt_boxes, pred_centerness95 )96 97 @torch.no_grad()98 def _match_anchors(self, gt_boxes: Boxes, anchors: List[Boxes]):99 """100 Match ground-truth boxes to a set of multi-level anchors.101 102 Args:103 gt_boxes: Ground-truth boxes from instances of an image.104 anchors: List of anchors for each feature map (of different scales).105 106 Returns:107 torch.Tensor108 A tensor of shape `(M, R)`, given `M` ground-truth boxes and total109 `R` anchor points from all feature levels, indicating the quality110 of match between m-th box and r-th anchor. Higher value indicates111 better match.112 """113 # Naming convention: (M = ground-truth boxes, R = anchor points)114 # Anchor points are represented as square boxes of size = stride.115 num_anchors_per_level = [len(x) for x in anchors]116 anchors = Boxes.cat(anchors) # (R, 4)117 anchor_centers = anchors.get_centers() # (R, 2)118 anchor_sizes = anchors.tensor[:, 2] - anchors.tensor[:, 0] # (R, )119 120 lower_bound = anchor_sizes * 4121 lower_bound[: num_anchors_per_level[0]] = 0122 upper_bound = anchor_sizes * 8123 upper_bound[-num_anchors_per_level[-1] :] = float("inf")124 125 gt_centers = gt_boxes.get_centers()126 127 # FCOS with center sampling: anchor point must be close enough to128 # ground-truth box center.129 center_dists = (anchor_centers[None, :, :] - gt_centers[:, None, :]).abs_()130 sampling_regions = self.center_sampling_radius * anchor_sizes[None, :]131 132 match_quality_matrix = center_dists.max(dim=2).values < sampling_regions133 134 pairwise_dist = pairwise_point_box_distance(anchor_centers, gt_boxes)135 pairwise_dist = pairwise_dist.permute(1, 0, 2) # (M, R, 4)136 137 # The original FCOS anchor matching rule: anchor point must be inside GT.138 match_quality_matrix &= pairwise_dist.min(dim=2).values > 0139 140 # Multilevel anchor matching in FCOS: each anchor is only responsible141 # for certain scale range.142 pairwise_dist = pairwise_dist.max(dim=2).values143 match_quality_matrix &= (pairwise_dist > lower_bound[None, :]) & (144 pairwise_dist < upper_bound[None, :]145 )146 # Match the GT box with minimum area, if there are multiple GT matches.147 gt_areas = gt_boxes.area() # (M, )148 149 match_quality_matrix = match_quality_matrix.to(torch.float32)150 match_quality_matrix *= 1e8 - gt_areas[:, None]151 return match_quality_matrix # (M, R)152 153 @torch.no_grad()154 def label_anchors(self, anchors: List[Boxes], gt_instances: List[Instances]):155 """156 Same interface as :meth:`RetinaNet.label_anchors`, but implemented with FCOS157 anchor matching rule.158 159 Unlike RetinaNet, there are no ignored anchors.160 """161 162 gt_labels, matched_gt_boxes = [], []163 164 for inst in gt_instances:165 if len(inst) > 0:166 match_quality_matrix = self._match_anchors(inst.gt_boxes, anchors)167 168 # Find matched ground-truth box per anchor. Un-matched anchors are169 # assigned -1. This is equivalent to using an anchor matcher as used170 # in R-CNN/RetinaNet: `Matcher(thresholds=[1e-5], labels=[0, 1])`171 match_quality, matched_idxs = match_quality_matrix.max(dim=0)172 matched_idxs[match_quality < 1e-5] = -1173 174 matched_gt_boxes_i = inst.gt_boxes.tensor[matched_idxs.clip(min=0)]175 gt_labels_i = inst.gt_classes[matched_idxs.clip(min=0)]176 177 # Anchors with matched_idxs = -1 are labeled background.178 gt_labels_i[matched_idxs < 0] = self.num_classes179 else:180 matched_gt_boxes_i = torch.zeros_like(Boxes.cat(anchors).tensor)181 gt_labels_i = torch.full(182 (len(matched_gt_boxes_i),),183 fill_value=self.num_classes,184 dtype=torch.long,185 device=matched_gt_boxes_i.device,186 )187 188 gt_labels.append(gt_labels_i)189 matched_gt_boxes.append(matched_gt_boxes_i)190 191 return gt_labels, matched_gt_boxes192 193 def losses(194 self, anchors, pred_logits, gt_labels, pred_anchor_deltas, gt_boxes, pred_centerness195 ):196 """197 This method is almost identical to :meth:`RetinaNet.losses`, with an extra198 "loss_centerness" in the returned dict.199 """200 num_images = len(gt_labels)201 gt_labels = torch.stack(gt_labels) # (M, R)202 203 pos_mask = (gt_labels >= 0) & (gt_labels != self.num_classes)204 num_pos_anchors = pos_mask.sum().item()205 get_event_storage().put_scalar("num_pos_anchors", num_pos_anchors / num_images)206 normalizer = self._ema_update("loss_normalizer", max(num_pos_anchors, 1), 300)207 208 # classification and regression loss209 gt_labels_target = F.one_hot(gt_labels, num_classes=self.num_classes + 1)[210 :, :, :-1211 ] # no loss for the last (background) class212 loss_cls = sigmoid_focal_loss_jit(213 torch.cat(pred_logits, dim=1),214 gt_labels_target.to(pred_logits[0].dtype),215 alpha=self.focal_loss_alpha,216 gamma=self.focal_loss_gamma,217 reduction="sum",218 )219 220 loss_box_reg = _dense_box_regression_loss(221 anchors,222 self.box2box_transform,223 pred_anchor_deltas,224 gt_boxes,225 pos_mask,226 box_reg_loss_type="giou",227 )228 229 ctrness_targets = self.compute_ctrness_targets(anchors, gt_boxes) # (M, R)230 pred_centerness = torch.cat(pred_centerness, dim=1).squeeze(dim=2) # (M, R)231 ctrness_loss = F.binary_cross_entropy_with_logits(232 pred_centerness[pos_mask], ctrness_targets[pos_mask], reduction="sum"233 )234 return {235 "loss_fcos_cls": loss_cls / normalizer,236 "loss_fcos_loc": loss_box_reg / normalizer,237 "loss_fcos_ctr": ctrness_loss / normalizer,238 }239 240 def compute_ctrness_targets(self, anchors: List[Boxes], gt_boxes: List[torch.Tensor]):241 anchors = Boxes.cat(anchors).tensor # Rx4242 reg_targets = [self.box2box_transform.get_deltas(anchors, m) for m in gt_boxes]243 reg_targets = torch.stack(reg_targets, dim=0) # NxRx4244 if len(reg_targets) == 0:245 return reg_targets.new_zeros(len(reg_targets))246 left_right = reg_targets[:, :, [0, 2]]247 top_bottom = reg_targets[:, :, [1, 3]]248 ctrness = (left_right.min(dim=-1)[0] / left_right.max(dim=-1)[0]) * (249 top_bottom.min(dim=-1)[0] / top_bottom.max(dim=-1)[0]250 )251 return torch.sqrt(ctrness)252 253 def forward_inference(254 self,255 images: ImageList,256 features: List[torch.Tensor],257 predictions: List[List[torch.Tensor]],258 ):259 pred_logits, pred_anchor_deltas, pred_centerness = self._transpose_dense_predictions(260 predictions, [self.num_classes, 4, 1]261 )262 anchors = self.anchor_generator(features)263 264 results: List[Instances] = []265 for img_idx, image_size in enumerate(images.image_sizes):266 scores_per_image = [267 # Multiply and sqrt centerness & classification scores268 # (See eqn. 4 in https://arxiv.org/abs/2006.09214)269 torch.sqrt(x[img_idx].sigmoid_() * y[img_idx].sigmoid_())270 for x, y in zip(pred_logits, pred_centerness)271 ]272 deltas_per_image = [x[img_idx] for x in pred_anchor_deltas]273 results_per_image = self.inference_single_image(274 anchors, scores_per_image, deltas_per_image, image_size275 )276 results.append(results_per_image)277 return results278 279 def inference_single_image(280 self,281 anchors: List[Boxes],282 box_cls: List[torch.Tensor],283 box_delta: List[torch.Tensor],284 image_size: Tuple[int, int],285 ):286 """287 Identical to :meth:`RetinaNet.inference_single_image.288 """289 pred = self._decode_multi_level_predictions(290 anchors,291 box_cls,292 box_delta,293 self.test_score_thresh,294 self.test_topk_candidates,295 image_size,296 )297 keep = batched_nms(298 pred.pred_boxes.tensor, pred.scores, pred.pred_classes, self.test_nms_thresh299 )300 return pred[keep[: self.max_detections_per_image]]301 302 303class FCOSHead(RetinaNetHead):304 """305 The head used in :paper:`fcos`. It adds an additional centerness306 prediction branch on top of :class:`RetinaNetHead`.307 """308 309 def __init__(self, *, input_shape: List[ShapeSpec], conv_dims: List[int], **kwargs):310 super().__init__(input_shape=input_shape, conv_dims=conv_dims, num_anchors=1, **kwargs)311 # Unlike original FCOS, we do not add an additional learnable scale layer312 # because it's found to have no benefits after normalizing regression targets by stride.313 self._num_features = len(input_shape)314 self.ctrness = nn.Conv2d(conv_dims[-1], 1, kernel_size=3, stride=1, padding=1)315 torch.nn.init.normal_(self.ctrness.weight, std=0.01)316 torch.nn.init.constant_(self.ctrness.bias, 0)317 318 def forward(self, features):319 assert len(features) == self._num_features320 logits = []321 bbox_reg = []322 ctrness = []323 for feature in features:324 logits.append(self.cls_score(self.cls_subnet(feature)))325 bbox_feature = self.bbox_subnet(feature)326 bbox_reg.append(self.bbox_pred(bbox_feature))327 ctrness.append(self.ctrness(bbox_feature))328 return logits, bbox_reg, ctrness329 