Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import itertools3import logging4from typing import Dict, List5import torch6 7from detectron2.config import configurable8from detectron2.layers import ShapeSpec, batched_nms_rotated, cat9from detectron2.structures import Instances, RotatedBoxes, pairwise_iou_rotated10from detectron2.utils.memory import retry_if_cuda_oom11 12from ..box_regression import Box2BoxTransformRotated13from .build import PROPOSAL_GENERATOR_REGISTRY14from .proposal_utils import _is_tracing15from .rpn import RPN16 17logger = logging.getLogger(__name__)18 19 20def find_top_rrpn_proposals(21 proposals,22 pred_objectness_logits,23 image_sizes,24 nms_thresh,25 pre_nms_topk,26 post_nms_topk,27 min_box_size,28 training,29):30 """31 For each feature map, select the `pre_nms_topk` highest scoring proposals,32 apply NMS, clip proposals, and remove small boxes. Return the `post_nms_topk`33 highest scoring proposals among all the feature maps if `training` is True,34 otherwise, returns the highest `post_nms_topk` scoring proposals for each35 feature map.36 37 Args:38 proposals (list[Tensor]): A list of L tensors. Tensor i has shape (N, Hi*Wi*A, 5).39 All proposal predictions on the feature maps.40 pred_objectness_logits (list[Tensor]): A list of L tensors. Tensor i has shape (N, Hi*Wi*A).41 image_sizes (list[tuple]): sizes (h, w) for each image42 nms_thresh (float): IoU threshold to use for NMS43 pre_nms_topk (int): number of top k scoring proposals to keep before applying NMS.44 When RRPN is run on multiple feature maps (as in FPN) this number is per45 feature map.46 post_nms_topk (int): number of top k scoring proposals to keep after applying NMS.47 When RRPN is run on multiple feature maps (as in FPN) this number is total,48 over all feature maps.49 min_box_size(float): minimum proposal box side length in pixels (absolute units wrt50 input images).51 training (bool): True if proposals are to be used in training, otherwise False.52 This arg exists only to support a legacy bug; look for the "NB: Legacy bug ..."53 comment.54 55 Returns:56 proposals (list[Instances]): list of N Instances. The i-th Instances57 stores post_nms_topk object proposals for image i.58 """59 num_images = len(image_sizes)60 device = proposals[0].device61 62 # 1. Select top-k anchor for every level and every image63 topk_scores = [] # #lvl Tensor, each of shape N x topk64 topk_proposals = []65 level_ids = [] # #lvl Tensor, each of shape (topk,)66 batch_idx = torch.arange(num_images, device=device)67 for level_id, proposals_i, logits_i in zip(68 itertools.count(), proposals, pred_objectness_logits69 ):70 Hi_Wi_A = logits_i.shape[1]71 if isinstance(Hi_Wi_A, torch.Tensor): # it's a tensor in tracing72 num_proposals_i = torch.clamp(Hi_Wi_A, max=pre_nms_topk)73 else:74 num_proposals_i = min(Hi_Wi_A, pre_nms_topk)75 76 topk_scores_i, topk_idx = logits_i.topk(num_proposals_i, dim=1)77 78 # each is N x topk79 topk_proposals_i = proposals_i[batch_idx[:, None], topk_idx] # N x topk x 580 81 topk_proposals.append(topk_proposals_i)82 topk_scores.append(topk_scores_i)83 level_ids.append(torch.full((num_proposals_i,), level_id, dtype=torch.int64, device=device))84 85 # 2. Concat all levels together86 topk_scores = cat(topk_scores, dim=1)87 topk_proposals = cat(topk_proposals, dim=1)88 level_ids = cat(level_ids, dim=0)89 90 # 3. For each image, run a per-level NMS, and choose topk results.91 results = []92 for n, image_size in enumerate(image_sizes):93 boxes = RotatedBoxes(topk_proposals[n])94 scores_per_img = topk_scores[n]95 lvl = level_ids96 97 valid_mask = torch.isfinite(boxes.tensor).all(dim=1) & torch.isfinite(scores_per_img)98 if not valid_mask.all():99 if training:100 raise FloatingPointError(101 "Predicted boxes or scores contain Inf/NaN. Training has diverged."102 )103 boxes = boxes[valid_mask]104 scores_per_img = scores_per_img[valid_mask]105 lvl = lvl[valid_mask]106 boxes.clip(image_size)107 108 # filter empty boxes109 keep = boxes.nonempty(threshold=min_box_size)110 if _is_tracing() or keep.sum().item() != len(boxes):111 boxes, scores_per_img, lvl = (boxes[keep], scores_per_img[keep], lvl[keep])112 113 keep = batched_nms_rotated(boxes.tensor, scores_per_img, lvl, nms_thresh)114 # In Detectron1, there was different behavior during training vs. testing.115 # (https://github.com/facebookresearch/Detectron/issues/459)116 # During training, topk is over the proposals from *all* images in the training batch.117 # During testing, it is over the proposals for each image separately.118 # As a result, the training behavior becomes batch-dependent,119 # and the configuration "POST_NMS_TOPK_TRAIN" end up relying on the batch size.120 # This bug is addressed in Detectron2 to make the behavior independent of batch size.121 keep = keep[:post_nms_topk]122 123 res = Instances(image_size)124 res.proposal_boxes = boxes[keep]125 res.objectness_logits = scores_per_img[keep]126 results.append(res)127 return results128 129 130@PROPOSAL_GENERATOR_REGISTRY.register()131class RRPN(RPN):132 """133 Rotated Region Proposal Network described in :paper:`RRPN`.134 """135 136 @configurable137 def __init__(self, *args, **kwargs):138 super().__init__(*args, **kwargs)139 if self.anchor_boundary_thresh >= 0:140 raise NotImplementedError(141 "anchor_boundary_thresh is a legacy option not implemented for RRPN."142 )143 144 @classmethod145 def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec]):146 ret = super().from_config(cfg, input_shape)147 ret["box2box_transform"] = Box2BoxTransformRotated(weights=cfg.MODEL.RPN.BBOX_REG_WEIGHTS)148 return ret149 150 @torch.no_grad()151 def label_and_sample_anchors(self, anchors: List[RotatedBoxes], gt_instances: List[Instances]):152 """153 Args:154 anchors (list[RotatedBoxes]): anchors for each feature map.155 gt_instances: the ground-truth instances for each image.156 157 Returns:158 list[Tensor]:159 List of #img tensors. i-th element is a vector of labels whose length is160 the total number of anchors across feature maps. Label values are in {-1, 0, 1},161 with meanings: -1 = ignore; 0 = negative class; 1 = positive class.162 list[Tensor]:163 i-th element is a Nx5 tensor, where N is the total number of anchors across164 feature maps. The values are the matched gt boxes for each anchor.165 Values are undefined for those anchors not labeled as 1.166 """167 anchors = RotatedBoxes.cat(anchors)168 169 gt_boxes = [x.gt_boxes for x in gt_instances]170 del gt_instances171 172 gt_labels = []173 matched_gt_boxes = []174 for gt_boxes_i in gt_boxes:175 """176 gt_boxes_i: ground-truth boxes for i-th image177 """178 match_quality_matrix = retry_if_cuda_oom(pairwise_iou_rotated)(gt_boxes_i, anchors)179 matched_idxs, gt_labels_i = retry_if_cuda_oom(self.anchor_matcher)(match_quality_matrix)180 # Matching is memory-expensive and may result in CPU tensors. But the result is small181 gt_labels_i = gt_labels_i.to(device=gt_boxes_i.device)182 183 # A vector of labels (-1, 0, 1) for each anchor184 gt_labels_i = self._subsample_labels(gt_labels_i)185 186 if len(gt_boxes_i) == 0:187 # These values won't be used anyway since the anchor is labeled as background188 matched_gt_boxes_i = torch.zeros_like(anchors.tensor)189 else:190 # TODO wasted indexing computation for ignored boxes191 matched_gt_boxes_i = gt_boxes_i[matched_idxs].tensor192 193 gt_labels.append(gt_labels_i) # N,AHW194 matched_gt_boxes.append(matched_gt_boxes_i)195 return gt_labels, matched_gt_boxes196 197 @torch.no_grad()198 def predict_proposals(self, anchors, pred_objectness_logits, pred_anchor_deltas, image_sizes):199 pred_proposals = self._decode_proposals(anchors, pred_anchor_deltas)200 return find_top_rrpn_proposals(201 pred_proposals,202 pred_objectness_logits,203 image_sizes,204 self.nms_thresh,205 self.pre_nms_topk[self.training],206 self.post_nms_topk[self.training],207 self.min_box_size,208 self.training,209 )210 