Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
trident_rcnn.py117 linesDownload Raw Back to tridentnet
1# Copyright (c) Facebook, Inc. and its affiliates.2from detectron2.layers import batched_nms3from detectron2.modeling import ROI_HEADS_REGISTRY, StandardROIHeads4from detectron2.modeling.roi_heads.roi_heads import Res5ROIHeads5from detectron2.structures import Instances6 7 8def merge_branch_instances(instances, num_branch, nms_thresh, topk_per_image):9    """10    Merge detection results from different branches of TridentNet.11    Return detection results by applying non-maximum suppression (NMS) on bounding boxes12    and keep the unsuppressed boxes and other instances (e.g mask) if any.13 14    Args:15        instances (list[Instances]): A list of N * num_branch instances that store detection16            results. Contain N images and each image has num_branch instances.17        num_branch (int): Number of branches used for merging detection results for each image.18        nms_thresh (float):  The threshold to use for box non-maximum suppression. Value in [0, 1].19        topk_per_image (int): The number of top scoring detections to return. Set < 0 to return20            all detections.21 22    Returns:23        results: (list[Instances]): A list of N instances, one for each image in the batch,24            that stores the topk most confidence detections after merging results from multiple25            branches.26    """27    if num_branch == 1:28        return instances29 30    batch_size = len(instances) // num_branch31    results = []32    for i in range(batch_size):33        instance = Instances.cat([instances[i + batch_size * j] for j in range(num_branch)])34 35        # Apply per-class NMS36        keep = batched_nms(37            instance.pred_boxes.tensor, instance.scores, instance.pred_classes, nms_thresh38        )39        keep = keep[:topk_per_image]40        result = instance[keep]41 42        results.append(result)43 44    return results45 46 47@ROI_HEADS_REGISTRY.register()48class TridentRes5ROIHeads(Res5ROIHeads):49    """50    The TridentNet ROIHeads in a typical "C4" R-CNN model.51    See :class:`Res5ROIHeads`.52    """53 54    def __init__(self, cfg, input_shape):55        super().__init__(cfg, input_shape)56 57        self.num_branch = cfg.MODEL.TRIDENT.NUM_BRANCH58        self.trident_fast = cfg.MODEL.TRIDENT.TEST_BRANCH_IDX != -159 60    def forward(self, images, features, proposals, targets=None):61        """62        See :class:`Res5ROIHeads.forward`.63        """64        num_branch = self.num_branch if self.training or not self.trident_fast else 165        all_targets = targets * num_branch if targets is not None else None66        pred_instances, losses = super().forward(images, features, proposals, all_targets)67        del images, all_targets, targets68 69        if self.training:70            return pred_instances, losses71        else:72            pred_instances = merge_branch_instances(73                pred_instances,74                num_branch,75                self.box_predictor.test_nms_thresh,76                self.box_predictor.test_topk_per_image,77            )78 79            return pred_instances, {}80 81 82@ROI_HEADS_REGISTRY.register()83class TridentStandardROIHeads(StandardROIHeads):84    """85    The `StandardROIHeads` for TridentNet.86    See :class:`StandardROIHeads`.87    """88 89    def __init__(self, cfg, input_shape):90        super(TridentStandardROIHeads, self).__init__(cfg, input_shape)91 92        self.num_branch = cfg.MODEL.TRIDENT.NUM_BRANCH93        self.trident_fast = cfg.MODEL.TRIDENT.TEST_BRANCH_IDX != -194 95    def forward(self, images, features, proposals, targets=None):96        """97        See :class:`Res5ROIHeads.forward`.98        """99        # Use 1 branch if using trident_fast during inference.100        num_branch = self.num_branch if self.training or not self.trident_fast else 1101        # Duplicate targets for all branches in TridentNet.102        all_targets = targets * num_branch if targets is not None else None103        pred_instances, losses = super().forward(images, features, proposals, all_targets)104        del images, all_targets, targets105 106        if self.training:107            return pred_instances, losses108        else:109            pred_instances = merge_branch_instances(110                pred_instances,111                num_branch,112                self.box_predictor.test_nms_thresh,113                self.box_predictor.test_topk_per_image,114            )115 116            return pred_instances, {}117