Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
1# Copyright (c) Facebook, Inc. and its affiliates.2from typing import Dict, List, Optional, Tuple, Union3import torch4import torch.nn.functional as F5from torch import nn6 7from detectron2.config import configurable8from detectron2.layers import Conv2d, ShapeSpec, cat9from detectron2.structures import Boxes, ImageList, Instances, pairwise_iou10from detectron2.utils.events import get_event_storage11from detectron2.utils.memory import retry_if_cuda_oom12from detectron2.utils.registry import Registry13 14from ..anchor_generator import build_anchor_generator15from ..box_regression import Box2BoxTransform, _dense_box_regression_loss16from ..matcher import Matcher17from ..sampling import subsample_labels18from .build import PROPOSAL_GENERATOR_REGISTRY19from .proposal_utils import find_top_rpn_proposals20 21RPN_HEAD_REGISTRY = Registry("RPN_HEAD")22RPN_HEAD_REGISTRY.__doc__ = """23Registry for RPN heads, which take feature maps and perform24objectness classification and bounding box regression for anchors.25 26The registered object will be called with `obj(cfg, input_shape)`.27The call should return a `nn.Module` object.28"""29 30 31"""32Shape shorthand in this module:33 34    N: number of images in the minibatch35    L: number of feature maps per image on which RPN is run36    A: number of cell anchors (must be the same for all feature maps)37    Hi, Wi: height and width of the i-th feature map38    B: size of the box parameterization39 40Naming convention:41 42    objectness: refers to the binary classification of an anchor as object vs. not object.43 44    deltas: refers to the 4-d (dx, dy, dw, dh) deltas that parameterize the box2box45    transform (see :class:`box_regression.Box2BoxTransform`), or 5d for rotated boxes.46 47    pred_objectness_logits: predicted objectness scores in [-inf, +inf]; use48        sigmoid(pred_objectness_logits) to estimate P(object).49 50    gt_labels: ground-truth binary classification labels for objectness51 52    pred_anchor_deltas: predicted box2box transform deltas53 54    gt_anchor_deltas: ground-truth box2box transform deltas55"""56 57 58def build_rpn_head(cfg, input_shape):59    """60    Build an RPN head defined by `cfg.MODEL.RPN.HEAD_NAME`.61    """62    name = cfg.MODEL.RPN.HEAD_NAME63    return RPN_HEAD_REGISTRY.get(name)(cfg, input_shape)64 65 66@RPN_HEAD_REGISTRY.register()67class StandardRPNHead(nn.Module):68    """69    Standard RPN classification and regression heads described in :paper:`Faster R-CNN`.70    Uses a 3x3 conv to produce a shared hidden state from which one 1x1 conv predicts71    objectness logits for each anchor and a second 1x1 conv predicts bounding-box deltas72    specifying how to deform each anchor into an object proposal.73    """74 75    @configurable76    def __init__(77        self, *, in_channels: int, num_anchors: int, box_dim: int = 4, conv_dims: List[int] = (-1,)78    ):79        """80        NOTE: this interface is experimental.81 82        Args:83            in_channels (int): number of input feature channels. When using multiple84                input features, they must have the same number of channels.85            num_anchors (int): number of anchors to predict for *each spatial position*86                on the feature map. The total number of anchors for each87                feature map will be `num_anchors * H * W`.88            box_dim (int): dimension of a box, which is also the number of box regression89                predictions to make for each anchor. An axis aligned box has90                box_dim=4, while a rotated box has box_dim=5.91            conv_dims (list[int]): a list of integers representing the output channels92                of N conv layers. Set it to -1 to use the same number of output channels93                as input channels.94        """95        super().__init__()96        cur_channels = in_channels97        # Keeping the old variable names and structure for backwards compatiblity.98        # Otherwise the old checkpoints will fail to load.99        if len(conv_dims) == 1:100            out_channels = cur_channels if conv_dims[0] == -1 else conv_dims[0]101            # 3x3 conv for the hidden representation102            self.conv = self._get_rpn_conv(cur_channels, out_channels)103            cur_channels = out_channels104        else:105            self.conv = nn.Sequential()106            for k, conv_dim in enumerate(conv_dims):107                out_channels = cur_channels if conv_dim == -1 else conv_dim108                if out_channels <= 0:109                    raise ValueError(110                        f"Conv output channels should be greater than 0. Got {out_channels}"111                    )112                conv = self._get_rpn_conv(cur_channels, out_channels)113                self.conv.add_module(f"conv{k}", conv)114                cur_channels = out_channels115        # 1x1 conv for predicting objectness logits116        self.objectness_logits = nn.Conv2d(cur_channels, num_anchors, kernel_size=1, stride=1)117        # 1x1 conv for predicting box2box transform deltas118        self.anchor_deltas = nn.Conv2d(cur_channels, num_anchors * box_dim, kernel_size=1, stride=1)119 120        # Keeping the order of weights initialization same for backwards compatiblility.121        for layer in self.modules():122            if isinstance(layer, nn.Conv2d):123                nn.init.normal_(layer.weight, std=0.01)124                nn.init.constant_(layer.bias, 0)125 126    def _get_rpn_conv(self, in_channels, out_channels):127        return Conv2d(128            in_channels,129            out_channels,130            kernel_size=3,131            stride=1,132            padding=1,133            activation=nn.ReLU(),134        )135 136    @classmethod137    def from_config(cls, cfg, input_shape):138        # Standard RPN is shared across levels:139        in_channels = [s.channels for s in input_shape]140        assert len(set(in_channels)) == 1, "Each level must have the same channel!"141        in_channels = in_channels[0]142 143        # RPNHead should take the same input as anchor generator144        # NOTE: it assumes that creating an anchor generator does not have unwanted side effect.145        anchor_generator = build_anchor_generator(cfg, input_shape)146        num_anchors = anchor_generator.num_anchors147        box_dim = anchor_generator.box_dim148        assert (149            len(set(num_anchors)) == 1150        ), "Each level must have the same number of anchors per spatial position"151        return {152            "in_channels": in_channels,153            "num_anchors": num_anchors[0],154            "box_dim": box_dim,155            "conv_dims": cfg.MODEL.RPN.CONV_DIMS,156        }157 158    def forward(self, features: List[torch.Tensor]):159        """160        Args:161            features (list[Tensor]): list of feature maps162 163        Returns:164            list[Tensor]: A list of L elements.165                Element i is a tensor of shape (N, A, Hi, Wi) representing166                the predicted objectness logits for all anchors. A is the number of cell anchors.167            list[Tensor]: A list of L elements. Element i is a tensor of shape168                (N, A*box_dim, Hi, Wi) representing the predicted "deltas" used to transform anchors169                to proposals.170        """171        pred_objectness_logits = []172        pred_anchor_deltas = []173        for x in features:174            t = self.conv(x)175            pred_objectness_logits.append(self.objectness_logits(t))176            pred_anchor_deltas.append(self.anchor_deltas(t))177        return pred_objectness_logits, pred_anchor_deltas178 179 180@PROPOSAL_GENERATOR_REGISTRY.register()181class RPN(nn.Module):182    """183    Region Proposal Network, introduced by :paper:`Faster R-CNN`.184    """185 186    @configurable187    def __init__(188        self,189        *,190        in_features: List[str],191        head: nn.Module,192        anchor_generator: nn.Module,193        anchor_matcher: Matcher,194        box2box_transform: Box2BoxTransform,195        batch_size_per_image: int,196        positive_fraction: float,197        pre_nms_topk: Tuple[float, float],198        post_nms_topk: Tuple[float, float],199        nms_thresh: float = 0.7,200        min_box_size: float = 0.0,201        anchor_boundary_thresh: float = -1.0,202        loss_weight: Union[float, Dict[str, float]] = 1.0,203        box_reg_loss_type: str = "smooth_l1",204        smooth_l1_beta: float = 0.0,205    ):206        """207        NOTE: this interface is experimental.208 209        Args:210            in_features (list[str]): list of names of input features to use211            head (nn.Module): a module that predicts logits and regression deltas212                for each level from a list of per-level features213            anchor_generator (nn.Module): a module that creates anchors from a214                list of features. Usually an instance of :class:`AnchorGenerator`215            anchor_matcher (Matcher): label the anchors by matching them with ground truth.216            box2box_transform (Box2BoxTransform): defines the transform from anchors boxes to217                instance boxes218            batch_size_per_image (int): number of anchors per image to sample for training219            positive_fraction (float): fraction of foreground anchors to sample for training220            pre_nms_topk (tuple[float]): (train, test) that represents the221                number of top k proposals to select before NMS, in222                training and testing.223            post_nms_topk (tuple[float]): (train, test) that represents the224                number of top k proposals to select after NMS, in225                training and testing.226            nms_thresh (float): NMS threshold used to de-duplicate the predicted proposals227            min_box_size (float): remove proposal boxes with any side smaller than this threshold,228                in the unit of input image pixels229            anchor_boundary_thresh (float): legacy option230            loss_weight (float|dict): weights to use for losses. Can be single float for weighting231                all rpn losses together, or a dict of individual weightings. Valid dict keys are:232                    "loss_rpn_cls" - applied to classification loss233                    "loss_rpn_loc" - applied to box regression loss234            box_reg_loss_type (str): Loss type to use. Supported losses: "smooth_l1", "giou".235            smooth_l1_beta (float): beta parameter for the smooth L1 regression loss. Default to236                use L1 loss. Only used when `box_reg_loss_type` is "smooth_l1"237        """238        super().__init__()239        self.in_features = in_features240        self.rpn_head = head241        self.anchor_generator = anchor_generator242        self.anchor_matcher = anchor_matcher243        self.box2box_transform = box2box_transform244        self.batch_size_per_image = batch_size_per_image245        self.positive_fraction = positive_fraction246        # Map from self.training state to train/test settings247        self.pre_nms_topk = {True: pre_nms_topk[0], False: pre_nms_topk[1]}248        self.post_nms_topk = {True: post_nms_topk[0], False: post_nms_topk[1]}249        self.nms_thresh = nms_thresh250        self.min_box_size = float(min_box_size)251        self.anchor_boundary_thresh = anchor_boundary_thresh252        if isinstance(loss_weight, float):253            loss_weight = {"loss_rpn_cls": loss_weight, "loss_rpn_loc": loss_weight}254        self.loss_weight = loss_weight255        self.box_reg_loss_type = box_reg_loss_type256        self.smooth_l1_beta = smooth_l1_beta257 258    @classmethod259    def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec]):260        in_features = cfg.MODEL.RPN.IN_FEATURES261        ret = {262            "in_features": in_features,263            "min_box_size": cfg.MODEL.PROPOSAL_GENERATOR.MIN_SIZE,264            "nms_thresh": cfg.MODEL.RPN.NMS_THRESH,265            "batch_size_per_image": cfg.MODEL.RPN.BATCH_SIZE_PER_IMAGE,266            "positive_fraction": cfg.MODEL.RPN.POSITIVE_FRACTION,267            "loss_weight": {268                "loss_rpn_cls": cfg.MODEL.RPN.LOSS_WEIGHT,269                "loss_rpn_loc": cfg.MODEL.RPN.BBOX_REG_LOSS_WEIGHT * cfg.MODEL.RPN.LOSS_WEIGHT,270            },271            "anchor_boundary_thresh": cfg.MODEL.RPN.BOUNDARY_THRESH,272            "box2box_transform": Box2BoxTransform(weights=cfg.MODEL.RPN.BBOX_REG_WEIGHTS),273            "box_reg_loss_type": cfg.MODEL.RPN.BBOX_REG_LOSS_TYPE,274            "smooth_l1_beta": cfg.MODEL.RPN.SMOOTH_L1_BETA,275        }276 277        ret["pre_nms_topk"] = (cfg.MODEL.RPN.PRE_NMS_TOPK_TRAIN, cfg.MODEL.RPN.PRE_NMS_TOPK_TEST)278        ret["post_nms_topk"] = (cfg.MODEL.RPN.POST_NMS_TOPK_TRAIN, cfg.MODEL.RPN.POST_NMS_TOPK_TEST)279 280        ret["anchor_generator"] = build_anchor_generator(cfg, [input_shape[f] for f in in_features])281        ret["anchor_matcher"] = Matcher(282            cfg.MODEL.RPN.IOU_THRESHOLDS, cfg.MODEL.RPN.IOU_LABELS, allow_low_quality_matches=True283        )284        ret["head"] = build_rpn_head(cfg, [input_shape[f] for f in in_features])285        return ret286 287    def _subsample_labels(self, label):288        """289        Randomly sample a subset of positive and negative examples, and overwrite290        the label vector to the ignore value (-1) for all elements that are not291        included in the sample.292 293        Args:294            labels (Tensor): a vector of -1, 0, 1. Will be modified in-place and returned.295        """296        pos_idx, neg_idx = subsample_labels(297            label, self.batch_size_per_image, self.positive_fraction, 0298        )299        # Fill with the ignore label (-1), then set positive and negative labels300        label.fill_(-1)301        label.scatter_(0, pos_idx, 1)302        label.scatter_(0, neg_idx, 0)303        return label304 305    @torch.jit.unused306    @torch.no_grad()307    def label_and_sample_anchors(308        self, anchors: List[Boxes], gt_instances: List[Instances]309    ) -> Tuple[List[torch.Tensor], List[torch.Tensor]]:310        """311        Args:312            anchors (list[Boxes]): anchors for each feature map.313            gt_instances: the ground-truth instances for each image.314 315        Returns:316            list[Tensor]:317                List of #img tensors. i-th element is a vector of labels whose length is318                the total number of anchors across all feature maps R = sum(Hi * Wi * A).319                Label values are in {-1, 0, 1}, with meanings: -1 = ignore; 0 = negative320                class; 1 = positive class.321            list[Tensor]:322                i-th element is a Rx4 tensor. The values are the matched gt boxes for each323                anchor. Values are undefined for those anchors not labeled as 1.324        """325        anchors = Boxes.cat(anchors)326 327        gt_boxes = [x.gt_boxes for x in gt_instances]328        image_sizes = [x.image_size for x in gt_instances]329        del gt_instances330 331        gt_labels = []332        matched_gt_boxes = []333        for image_size_i, gt_boxes_i in zip(image_sizes, gt_boxes):334            """335            image_size_i: (h, w) for the i-th image336            gt_boxes_i: ground-truth boxes for i-th image337            """338 339            match_quality_matrix = retry_if_cuda_oom(pairwise_iou)(gt_boxes_i, anchors)340            matched_idxs, gt_labels_i = retry_if_cuda_oom(self.anchor_matcher)(match_quality_matrix)341            # Matching is memory-expensive and may result in CPU tensors. But the result is small342            gt_labels_i = gt_labels_i.to(device=gt_boxes_i.device)343            del match_quality_matrix344 345            if self.anchor_boundary_thresh >= 0:346                # Discard anchors that go out of the boundaries of the image347                # NOTE: This is legacy functionality that is turned off by default in Detectron2348                anchors_inside_image = anchors.inside_box(image_size_i, self.anchor_boundary_thresh)349                gt_labels_i[~anchors_inside_image] = -1350 351            # A vector of labels (-1, 0, 1) for each anchor352            gt_labels_i = self._subsample_labels(gt_labels_i)353 354            if len(gt_boxes_i) == 0:355                # These values won't be used anyway since the anchor is labeled as background356                matched_gt_boxes_i = torch.zeros_like(anchors.tensor)357            else:358                # TODO wasted indexing computation for ignored boxes359                matched_gt_boxes_i = gt_boxes_i[matched_idxs].tensor360 361            gt_labels.append(gt_labels_i)  # N,AHW362            matched_gt_boxes.append(matched_gt_boxes_i)363        return gt_labels, matched_gt_boxes364 365    @torch.jit.unused366    def losses(367        self,368        anchors: List[Boxes],369        pred_objectness_logits: List[torch.Tensor],370        gt_labels: List[torch.Tensor],371        pred_anchor_deltas: List[torch.Tensor],372        gt_boxes: List[torch.Tensor],373    ) -> Dict[str, torch.Tensor]:374        """375        Return the losses from a set of RPN predictions and their associated ground-truth.376 377        Args:378            anchors (list[Boxes or RotatedBoxes]): anchors for each feature map, each379                has shape (Hi*Wi*A, B), where B is box dimension (4 or 5).380            pred_objectness_logits (list[Tensor]): A list of L elements.381                Element i is a tensor of shape (N, Hi*Wi*A) representing382                the predicted objectness logits for all anchors.383            gt_labels (list[Tensor]): Output of :meth:`label_and_sample_anchors`.384            pred_anchor_deltas (list[Tensor]): A list of L elements. Element i is a tensor of shape385                (N, Hi*Wi*A, 4 or 5) representing the predicted "deltas" used to transform anchors386                to proposals.387            gt_boxes (list[Tensor]): Output of :meth:`label_and_sample_anchors`.388 389        Returns:390            dict[loss name -> loss value]: A dict mapping from loss name to loss value.391                Loss names are: `loss_rpn_cls` for objectness classification and392                `loss_rpn_loc` for proposal localization.393        """394        num_images = len(gt_labels)395        gt_labels = torch.stack(gt_labels)  # (N, sum(Hi*Wi*Ai))396 397        # Log the number of positive/negative anchors per-image that's used in training398        pos_mask = gt_labels == 1399        num_pos_anchors = pos_mask.sum().item()400        num_neg_anchors = (gt_labels == 0).sum().item()401        storage = get_event_storage()402        storage.put_scalar("rpn/num_pos_anchors", num_pos_anchors / num_images)403        storage.put_scalar("rpn/num_neg_anchors", num_neg_anchors / num_images)404 405        localization_loss = _dense_box_regression_loss(406            anchors,407            self.box2box_transform,408            pred_anchor_deltas,409            gt_boxes,410            pos_mask,411            box_reg_loss_type=self.box_reg_loss_type,412            smooth_l1_beta=self.smooth_l1_beta,413        )414 415        valid_mask = gt_labels >= 0416        objectness_loss = F.binary_cross_entropy_with_logits(417            cat(pred_objectness_logits, dim=1)[valid_mask],418            gt_labels[valid_mask].to(torch.float32),419            reduction="sum",420        )421        normalizer = self.batch_size_per_image * num_images422        losses = {423            "loss_rpn_cls": objectness_loss / normalizer,424            # The original Faster R-CNN paper uses a slightly different normalizer425            # for loc loss. But it doesn't matter in practice426            "loss_rpn_loc": localization_loss / normalizer,427        }428        losses = {k: v * self.loss_weight.get(k, 1.0) for k, v in losses.items()}429        return losses430 431    def forward(432        self,433        images: ImageList,434        features: Dict[str, torch.Tensor],435        gt_instances: Optional[List[Instances]] = None,436    ):437        """438        Args:439            images (ImageList): input images of length `N`440            features (dict[str, Tensor]): input data as a mapping from feature441                map name to tensor. Axis 0 represents the number of images `N` in442                the input data; axes 1-3 are channels, height, and width, which may443                vary between feature maps (e.g., if a feature pyramid is used).444            gt_instances (list[Instances], optional): a length `N` list of `Instances`s.445                Each `Instances` stores ground-truth instances for the corresponding image.446 447        Returns:448            proposals: list[Instances]: contains fields "proposal_boxes", "objectness_logits"449            loss: dict[Tensor] or None450        """451        features = [features[f] for f in self.in_features]452        anchors = self.anchor_generator(features)453 454        pred_objectness_logits, pred_anchor_deltas = self.rpn_head(features)455        # Transpose the Hi*Wi*A dimension to the middle:456        pred_objectness_logits = [457            # (N, A, Hi, Wi) -> (N, Hi, Wi, A) -> (N, Hi*Wi*A)458            score.permute(0, 2, 3, 1).flatten(1)459            for score in pred_objectness_logits460        ]461        pred_anchor_deltas = [462            # (N, A*B, Hi, Wi) -> (N, A, B, Hi, Wi) -> (N, Hi, Wi, A, B) -> (N, Hi*Wi*A, B)463            x.view(x.shape[0], -1, self.anchor_generator.box_dim, x.shape[-2], x.shape[-1])464            .permute(0, 3, 4, 1, 2)465            .flatten(1, -2)466            for x in pred_anchor_deltas467        ]468 469        if self.training:470            assert gt_instances is not None, "RPN requires gt_instances in training!"471            gt_labels, gt_boxes = self.label_and_sample_anchors(anchors, gt_instances)472            losses = self.losses(473                anchors, pred_objectness_logits, gt_labels, pred_anchor_deltas, gt_boxes474            )475        else:476            losses = {}477        proposals = self.predict_proposals(478            anchors, pred_objectness_logits, pred_anchor_deltas, images.image_sizes479        )480        return proposals, losses481 482    def predict_proposals(483        self,484        anchors: List[Boxes],485        pred_objectness_logits: List[torch.Tensor],486        pred_anchor_deltas: List[torch.Tensor],487        image_sizes: List[Tuple[int, int]],488    ):489        """490        Decode all the predicted box regression deltas to proposals. Find the top proposals491        by applying NMS and removing boxes that are too small.492 493        Returns:494            proposals (list[Instances]): list of N Instances. The i-th Instances495                stores post_nms_topk object proposals for image i, sorted by their496                objectness score in descending order.497        """498        # The proposals are treated as fixed for joint training with roi heads.499        # This approach ignores the derivative w.r.t. the proposal boxes’ coordinates that500        # are also network responses.501        with torch.no_grad():502            pred_proposals = self._decode_proposals(anchors, pred_anchor_deltas)503            return find_top_rpn_proposals(504                pred_proposals,505                pred_objectness_logits,506                image_sizes,507                self.nms_thresh,508                self.pre_nms_topk[self.training],509                self.post_nms_topk[self.training],510                self.min_box_size,511                self.training,512            )513 514    def _decode_proposals(self, anchors: List[Boxes], pred_anchor_deltas: List[torch.Tensor]):515        """516        Transform anchors into proposals by applying the predicted anchor deltas.517 518        Returns:519            proposals (list[Tensor]): A list of L tensors. Tensor i has shape520                (N, Hi*Wi*A, B)521        """522        N = pred_anchor_deltas[0].shape[0]523        proposals = []524        # For each feature map525        for anchors_i, pred_anchor_deltas_i in zip(anchors, pred_anchor_deltas):526            B = anchors_i.tensor.size(1)527            pred_anchor_deltas_i = pred_anchor_deltas_i.reshape(-1, B)528            # Expand anchors to shape (N*Hi*Wi*A, B)529            anchors_i = anchors_i.tensor.unsqueeze(0).expand(N, -1, -1).reshape(-1, B)530            proposals_i = self.box2box_transform.apply_deltas(pred_anchor_deltas_i, anchors_i)531            # Append feature map proposals with shape (N, Hi*Wi*A, B)532            proposals.append(proposals_i.view(N, -1, B))533        return proposals534