Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import logging3import math4from typing import List, Tuple, Union5import torch6 7from detectron2.layers import batched_nms, cat, move_device_like8from detectron2.structures import Boxes, Instances9 10logger = logging.getLogger(__name__)11 12 13def _is_tracing():14 # (fixed in TORCH_VERSION >= 1.9)15 if torch.jit.is_scripting():16 # https://github.com/pytorch/pytorch/issues/4737917 return False18 else:19 return torch.jit.is_tracing()20 21 22def find_top_rpn_proposals(23 proposals: List[torch.Tensor],24 pred_objectness_logits: List[torch.Tensor],25 image_sizes: List[Tuple[int, int]],26 nms_thresh: float,27 pre_nms_topk: int,28 post_nms_topk: int,29 min_box_size: float,30 training: bool,31):32 """33 For each feature map, select the `pre_nms_topk` highest scoring proposals,34 apply NMS, clip proposals, and remove small boxes. Return the `post_nms_topk`35 highest scoring proposals among all the feature maps for each image.36 37 Args:38 proposals (list[Tensor]): A list of L tensors. Tensor i has shape (N, Hi*Wi*A, 4).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 RPN 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 RPN 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 units50 wrt 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 list[Instances]: list of N Instances. The i-th Instances57 stores post_nms_topk object proposals for image i, sorted by their58 objectness score in descending order.59 """60 num_images = len(image_sizes)61 device = (62 proposals[0].device63 if torch.jit.is_scripting()64 else ("cpu" if torch.jit.is_tracing() else proposals[0].device)65 )66 67 # 1. Select top-k anchor for every level and every image68 topk_scores = [] # #lvl Tensor, each of shape N x topk69 topk_proposals = []70 level_ids = [] # #lvl Tensor, each of shape (topk,)71 batch_idx = move_device_like(torch.arange(num_images, device=device), proposals[0])72 for level_id, (proposals_i, logits_i) in enumerate(zip(proposals, pred_objectness_logits)):73 Hi_Wi_A = logits_i.shape[1]74 if isinstance(Hi_Wi_A, torch.Tensor): # it's a tensor in tracing75 num_proposals_i = torch.clamp(Hi_Wi_A, max=pre_nms_topk)76 else:77 num_proposals_i = min(Hi_Wi_A, pre_nms_topk)78 79 topk_scores_i, topk_idx = logits_i.topk(num_proposals_i, dim=1)80 81 # each is N x topk82 topk_proposals_i = proposals_i[batch_idx[:, None], topk_idx] # N x topk x 483 84 topk_proposals.append(topk_proposals_i)85 topk_scores.append(topk_scores_i)86 level_ids.append(87 move_device_like(88 torch.full((num_proposals_i,), level_id, dtype=torch.int64, device=device),89 proposals[0],90 )91 )92 93 # 2. Concat all levels together94 topk_scores = cat(topk_scores, dim=1)95 topk_proposals = cat(topk_proposals, dim=1)96 level_ids = cat(level_ids, dim=0)97 98 # 3. For each image, run a per-level NMS, and choose topk results.99 results: List[Instances] = []100 for n, image_size in enumerate(image_sizes):101 boxes = Boxes(topk_proposals[n])102 scores_per_img = topk_scores[n]103 lvl = level_ids104 105 valid_mask = torch.isfinite(boxes.tensor).all(dim=1) & torch.isfinite(scores_per_img)106 if not valid_mask.all():107 if training:108 raise FloatingPointError(109 "Predicted boxes or scores contain Inf/NaN. Training has diverged."110 )111 boxes = boxes[valid_mask]112 scores_per_img = scores_per_img[valid_mask]113 lvl = lvl[valid_mask]114 boxes.clip(image_size)115 116 # filter empty boxes117 keep = boxes.nonempty(threshold=min_box_size)118 if _is_tracing() or keep.sum().item() != len(boxes):119 boxes, scores_per_img, lvl = boxes[keep], scores_per_img[keep], lvl[keep]120 121 keep = batched_nms(boxes.tensor, scores_per_img, lvl, nms_thresh)122 # In Detectron1, there was different behavior during training vs. testing.123 # (https://github.com/facebookresearch/Detectron/issues/459)124 # During training, topk is over the proposals from *all* images in the training batch.125 # During testing, it is over the proposals for each image separately.126 # As a result, the training behavior becomes batch-dependent,127 # and the configuration "POST_NMS_TOPK_TRAIN" end up relying on the batch size.128 # This bug is addressed in Detectron2 to make the behavior independent of batch size.129 keep = keep[:post_nms_topk] # keep is already sorted130 131 res = Instances(image_size)132 res.proposal_boxes = boxes[keep]133 res.objectness_logits = scores_per_img[keep]134 results.append(res)135 return results136 137 138def add_ground_truth_to_proposals(139 gt: Union[List[Instances], List[Boxes]], proposals: List[Instances]140) -> List[Instances]:141 """142 Call `add_ground_truth_to_proposals_single_image` for all images.143 144 Args:145 gt(Union[List[Instances], List[Boxes]): list of N elements. Element i is a Instances146 representing the ground-truth for image i.147 proposals (list[Instances]): list of N elements. Element i is a Instances148 representing the proposals for image i.149 150 Returns:151 list[Instances]: list of N Instances. Each is the proposals for the image,152 with field "proposal_boxes" and "objectness_logits".153 """154 assert gt is not None155 156 if len(proposals) != len(gt):157 raise ValueError("proposals and gt should have the same length as the number of images!")158 if len(proposals) == 0:159 return proposals160 161 return [162 add_ground_truth_to_proposals_single_image(gt_i, proposals_i)163 for gt_i, proposals_i in zip(gt, proposals)164 ]165 166 167def add_ground_truth_to_proposals_single_image(168 gt: Union[Instances, Boxes], proposals: Instances169) -> Instances:170 """171 Augment `proposals` with `gt`.172 173 Args:174 Same as `add_ground_truth_to_proposals`, but with gt and proposals175 per image.176 177 Returns:178 Same as `add_ground_truth_to_proposals`, but for only one image.179 """180 if isinstance(gt, Boxes):181 # convert Boxes to Instances182 gt = Instances(proposals.image_size, gt_boxes=gt)183 184 gt_boxes = gt.gt_boxes185 device = proposals.objectness_logits.device186 # Assign all ground-truth boxes an objectness logit corresponding to187 # P(object) = sigmoid(logit) =~ 1.188 gt_logit_value = math.log((1.0 - 1e-10) / (1 - (1.0 - 1e-10)))189 gt_logits = gt_logit_value * torch.ones(len(gt_boxes), device=device)190 191 # Concatenating gt_boxes with proposals requires them to have the same fields192 gt_proposal = Instances(proposals.image_size, **gt.get_fields())193 gt_proposal.proposal_boxes = gt_boxes194 gt_proposal.objectness_logits = gt_logits195 196 for key in proposals.get_fields().keys():197 assert gt_proposal.has(198 key199 ), "The attribute '{}' in `proposals` does not exist in `gt`".format(key)200 201 # NOTE: Instances.cat only use fields from the first item. Extra fields in latter items202 # will be thrown away.203 new_proposals = Instances.cat([proposals, gt_proposal])204 205 return new_proposals206 