Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
proposal_utils.py206 linesDownload Raw Back to proposal_generator
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