Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
fcos.py329 linesDownload Raw Back to meta_arch
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import logging4from typing import List, Optional, Tuple5import torch6from fvcore.nn import sigmoid_focal_loss_jit7from torch import nn8from torch.nn import functional as F9 10from detectron2.layers import ShapeSpec, batched_nms11from detectron2.structures import Boxes, ImageList, Instances, pairwise_point_box_distance12from detectron2.utils.events import get_event_storage13 14from ..anchor_generator import DefaultAnchorGenerator15from ..backbone import Backbone16from ..box_regression import Box2BoxTransformLinear, _dense_box_regression_loss17from .dense_detector import DenseDetector18from .retinanet import RetinaNetHead19 20__all__ = ["FCOS"]21 22logger = logging.getLogger(__name__)23 24 25class FCOS(DenseDetector):26    """27    Implement FCOS in :paper:`fcos`.28    """29 30    def __init__(31        self,32        *,33        backbone: Backbone,34        head: nn.Module,35        head_in_features: Optional[List[str]] = None,36        box2box_transform=None,37        num_classes,38        center_sampling_radius: float = 1.5,39        focal_loss_alpha=0.25,40        focal_loss_gamma=2.0,41        test_score_thresh=0.2,42        test_topk_candidates=1000,43        test_nms_thresh=0.6,44        max_detections_per_image=100,45        pixel_mean,46        pixel_std,47    ):48        """49        Args:50            center_sampling_radius: radius of the "center" of a groundtruth box,51                within which all anchor points are labeled positive.52            Other arguments mean the same as in :class:`RetinaNet`.53        """54        super().__init__(55            backbone, head, head_in_features, pixel_mean=pixel_mean, pixel_std=pixel_std56        )57 58        self.num_classes = num_classes59 60        # FCOS uses one anchor point per location.61        # We represent the anchor point by a box whose size equals the anchor stride.62        feature_shapes = backbone.output_shape()63        fpn_strides = [feature_shapes[k].stride for k in self.head_in_features]64        self.anchor_generator = DefaultAnchorGenerator(65            sizes=[[k] for k in fpn_strides], aspect_ratios=[1.0], strides=fpn_strides66        )67 68        # FCOS parameterizes box regression by a linear transform,69        # where predictions are normalized by anchor stride (equal to anchor size).70        if box2box_transform is None:71            box2box_transform = Box2BoxTransformLinear(normalize_by_size=True)72        self.box2box_transform = box2box_transform73 74        self.center_sampling_radius = float(center_sampling_radius)75 76        # Loss parameters:77        self.focal_loss_alpha = focal_loss_alpha78        self.focal_loss_gamma = focal_loss_gamma79 80        # Inference parameters:81        self.test_score_thresh = test_score_thresh82        self.test_topk_candidates = test_topk_candidates83        self.test_nms_thresh = test_nms_thresh84        self.max_detections_per_image = max_detections_per_image85 86    def forward_training(self, images, features, predictions, gt_instances):87        # Transpose the Hi*Wi*A dimension to the middle:88        pred_logits, pred_anchor_deltas, pred_centerness = self._transpose_dense_predictions(89            predictions, [self.num_classes, 4, 1]90        )91        anchors = self.anchor_generator(features)92        gt_labels, gt_boxes = self.label_anchors(anchors, gt_instances)93        return self.losses(94            anchors, pred_logits, gt_labels, pred_anchor_deltas, gt_boxes, pred_centerness95        )96 97    @torch.no_grad()98    def _match_anchors(self, gt_boxes: Boxes, anchors: List[Boxes]):99        """100        Match ground-truth boxes to a set of multi-level anchors.101 102        Args:103            gt_boxes: Ground-truth boxes from instances of an image.104            anchors: List of anchors for each feature map (of different scales).105 106        Returns:107            torch.Tensor108                A tensor of shape `(M, R)`, given `M` ground-truth boxes and total109                `R` anchor points from all feature levels, indicating the quality110                of match between m-th box and r-th anchor. Higher value indicates111                better match.112        """113        # Naming convention: (M = ground-truth boxes, R = anchor points)114        # Anchor points are represented as square boxes of size = stride.115        num_anchors_per_level = [len(x) for x in anchors]116        anchors = Boxes.cat(anchors)  # (R, 4)117        anchor_centers = anchors.get_centers()  # (R, 2)118        anchor_sizes = anchors.tensor[:, 2] - anchors.tensor[:, 0]  # (R, )119 120        lower_bound = anchor_sizes * 4121        lower_bound[: num_anchors_per_level[0]] = 0122        upper_bound = anchor_sizes * 8123        upper_bound[-num_anchors_per_level[-1] :] = float("inf")124 125        gt_centers = gt_boxes.get_centers()126 127        # FCOS with center sampling: anchor point must be close enough to128        # ground-truth box center.129        center_dists = (anchor_centers[None, :, :] - gt_centers[:, None, :]).abs_()130        sampling_regions = self.center_sampling_radius * anchor_sizes[None, :]131 132        match_quality_matrix = center_dists.max(dim=2).values < sampling_regions133 134        pairwise_dist = pairwise_point_box_distance(anchor_centers, gt_boxes)135        pairwise_dist = pairwise_dist.permute(1, 0, 2)  # (M, R, 4)136 137        # The original FCOS anchor matching rule: anchor point must be inside GT.138        match_quality_matrix &= pairwise_dist.min(dim=2).values > 0139 140        # Multilevel anchor matching in FCOS: each anchor is only responsible141        # for certain scale range.142        pairwise_dist = pairwise_dist.max(dim=2).values143        match_quality_matrix &= (pairwise_dist > lower_bound[None, :]) & (144            pairwise_dist < upper_bound[None, :]145        )146        # Match the GT box with minimum area, if there are multiple GT matches.147        gt_areas = gt_boxes.area()  # (M, )148 149        match_quality_matrix = match_quality_matrix.to(torch.float32)150        match_quality_matrix *= 1e8 - gt_areas[:, None]151        return match_quality_matrix  # (M, R)152 153    @torch.no_grad()154    def label_anchors(self, anchors: List[Boxes], gt_instances: List[Instances]):155        """156        Same interface as :meth:`RetinaNet.label_anchors`, but implemented with FCOS157        anchor matching rule.158 159        Unlike RetinaNet, there are no ignored anchors.160        """161 162        gt_labels, matched_gt_boxes = [], []163 164        for inst in gt_instances:165            if len(inst) > 0:166                match_quality_matrix = self._match_anchors(inst.gt_boxes, anchors)167 168                # Find matched ground-truth box per anchor. Un-matched anchors are169                # assigned -1. This is equivalent to using an anchor matcher as used170                # in R-CNN/RetinaNet: `Matcher(thresholds=[1e-5], labels=[0, 1])`171                match_quality, matched_idxs = match_quality_matrix.max(dim=0)172                matched_idxs[match_quality < 1e-5] = -1173 174                matched_gt_boxes_i = inst.gt_boxes.tensor[matched_idxs.clip(min=0)]175                gt_labels_i = inst.gt_classes[matched_idxs.clip(min=0)]176 177                # Anchors with matched_idxs = -1 are labeled background.178                gt_labels_i[matched_idxs < 0] = self.num_classes179            else:180                matched_gt_boxes_i = torch.zeros_like(Boxes.cat(anchors).tensor)181                gt_labels_i = torch.full(182                    (len(matched_gt_boxes_i),),183                    fill_value=self.num_classes,184                    dtype=torch.long,185                    device=matched_gt_boxes_i.device,186                )187 188            gt_labels.append(gt_labels_i)189            matched_gt_boxes.append(matched_gt_boxes_i)190 191        return gt_labels, matched_gt_boxes192 193    def losses(194        self, anchors, pred_logits, gt_labels, pred_anchor_deltas, gt_boxes, pred_centerness195    ):196        """197        This method is almost identical to :meth:`RetinaNet.losses`, with an extra198        "loss_centerness" in the returned dict.199        """200        num_images = len(gt_labels)201        gt_labels = torch.stack(gt_labels)  # (M, R)202 203        pos_mask = (gt_labels >= 0) & (gt_labels != self.num_classes)204        num_pos_anchors = pos_mask.sum().item()205        get_event_storage().put_scalar("num_pos_anchors", num_pos_anchors / num_images)206        normalizer = self._ema_update("loss_normalizer", max(num_pos_anchors, 1), 300)207 208        # classification and regression loss209        gt_labels_target = F.one_hot(gt_labels, num_classes=self.num_classes + 1)[210            :, :, :-1211        ]  # no loss for the last (background) class212        loss_cls = sigmoid_focal_loss_jit(213            torch.cat(pred_logits, dim=1),214            gt_labels_target.to(pred_logits[0].dtype),215            alpha=self.focal_loss_alpha,216            gamma=self.focal_loss_gamma,217            reduction="sum",218        )219 220        loss_box_reg = _dense_box_regression_loss(221            anchors,222            self.box2box_transform,223            pred_anchor_deltas,224            gt_boxes,225            pos_mask,226            box_reg_loss_type="giou",227        )228 229        ctrness_targets = self.compute_ctrness_targets(anchors, gt_boxes)  # (M, R)230        pred_centerness = torch.cat(pred_centerness, dim=1).squeeze(dim=2)  # (M, R)231        ctrness_loss = F.binary_cross_entropy_with_logits(232            pred_centerness[pos_mask], ctrness_targets[pos_mask], reduction="sum"233        )234        return {235            "loss_fcos_cls": loss_cls / normalizer,236            "loss_fcos_loc": loss_box_reg / normalizer,237            "loss_fcos_ctr": ctrness_loss / normalizer,238        }239 240    def compute_ctrness_targets(self, anchors: List[Boxes], gt_boxes: List[torch.Tensor]):241        anchors = Boxes.cat(anchors).tensor  # Rx4242        reg_targets = [self.box2box_transform.get_deltas(anchors, m) for m in gt_boxes]243        reg_targets = torch.stack(reg_targets, dim=0)  # NxRx4244        if len(reg_targets) == 0:245            return reg_targets.new_zeros(len(reg_targets))246        left_right = reg_targets[:, :, [0, 2]]247        top_bottom = reg_targets[:, :, [1, 3]]248        ctrness = (left_right.min(dim=-1)[0] / left_right.max(dim=-1)[0]) * (249            top_bottom.min(dim=-1)[0] / top_bottom.max(dim=-1)[0]250        )251        return torch.sqrt(ctrness)252 253    def forward_inference(254        self,255        images: ImageList,256        features: List[torch.Tensor],257        predictions: List[List[torch.Tensor]],258    ):259        pred_logits, pred_anchor_deltas, pred_centerness = self._transpose_dense_predictions(260            predictions, [self.num_classes, 4, 1]261        )262        anchors = self.anchor_generator(features)263 264        results: List[Instances] = []265        for img_idx, image_size in enumerate(images.image_sizes):266            scores_per_image = [267                # Multiply and sqrt centerness & classification scores268                # (See eqn. 4 in https://arxiv.org/abs/2006.09214)269                torch.sqrt(x[img_idx].sigmoid_() * y[img_idx].sigmoid_())270                for x, y in zip(pred_logits, pred_centerness)271            ]272            deltas_per_image = [x[img_idx] for x in pred_anchor_deltas]273            results_per_image = self.inference_single_image(274                anchors, scores_per_image, deltas_per_image, image_size275            )276            results.append(results_per_image)277        return results278 279    def inference_single_image(280        self,281        anchors: List[Boxes],282        box_cls: List[torch.Tensor],283        box_delta: List[torch.Tensor],284        image_size: Tuple[int, int],285    ):286        """287        Identical to :meth:`RetinaNet.inference_single_image.288        """289        pred = self._decode_multi_level_predictions(290            anchors,291            box_cls,292            box_delta,293            self.test_score_thresh,294            self.test_topk_candidates,295            image_size,296        )297        keep = batched_nms(298            pred.pred_boxes.tensor, pred.scores, pred.pred_classes, self.test_nms_thresh299        )300        return pred[keep[: self.max_detections_per_image]]301 302 303class FCOSHead(RetinaNetHead):304    """305    The head used in :paper:`fcos`. It adds an additional centerness306    prediction branch on top of :class:`RetinaNetHead`.307    """308 309    def __init__(self, *, input_shape: List[ShapeSpec], conv_dims: List[int], **kwargs):310        super().__init__(input_shape=input_shape, conv_dims=conv_dims, num_anchors=1, **kwargs)311        # Unlike original FCOS, we do not add an additional learnable scale layer312        # because it's found to have no benefits after normalizing regression targets by stride.313        self._num_features = len(input_shape)314        self.ctrness = nn.Conv2d(conv_dims[-1], 1, kernel_size=3, stride=1, padding=1)315        torch.nn.init.normal_(self.ctrness.weight, std=0.01)316        torch.nn.init.constant_(self.ctrness.bias, 0)317 318    def forward(self, features):319        assert len(features) == self._num_features320        logits = []321        bbox_reg = []322        ctrness = []323        for feature in features:324            logits.append(self.cls_score(self.cls_subnet(feature)))325            bbox_feature = self.bbox_subnet(feature)326            bbox_reg.append(self.bbox_pred(bbox_feature))327            ctrness.append(self.ctrness(bbox_feature))328        return logits, bbox_reg, ctrness329