Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
rcnn.py342 linesDownload Raw Back to meta_arch
1# Copyright (c) Facebook, Inc. and its affiliates.2import logging3import numpy as np4from typing import Dict, List, Optional, Tuple5import torch6from torch import nn7 8from detectron2.config import configurable9from detectron2.data.detection_utils import convert_image_to_rgb10from detectron2.layers import move_device_like11from detectron2.structures import ImageList, Instances12from detectron2.utils.events import get_event_storage13from detectron2.utils.logger import log_first_n14 15from ..backbone import Backbone, build_backbone16from ..postprocessing import detector_postprocess17from ..proposal_generator import build_proposal_generator18from ..roi_heads import build_roi_heads19from .build import META_ARCH_REGISTRY20 21__all__ = ["GeneralizedRCNN", "ProposalNetwork"]22 23 24@META_ARCH_REGISTRY.register()25class GeneralizedRCNN(nn.Module):26    """27    Generalized R-CNN. Any models that contains the following three components:28    1. Per-image feature extraction (aka backbone)29    2. Region proposal generation30    3. Per-region feature extraction and prediction31    """32 33    @configurable34    def __init__(35        self,36        *,37        backbone: Backbone,38        proposal_generator: nn.Module,39        roi_heads: nn.Module,40        pixel_mean: Tuple[float],41        pixel_std: Tuple[float],42        input_format: Optional[str] = None,43        vis_period: int = 0,44    ):45        """46        Args:47            backbone: a backbone module, must follow detectron2's backbone interface48            proposal_generator: a module that generates proposals using backbone features49            roi_heads: a ROI head that performs per-region computation50            pixel_mean, pixel_std: list or tuple with #channels element, representing51                the per-channel mean and std to be used to normalize the input image52            input_format: describe the meaning of channels of input. Needed by visualization53            vis_period: the period to run visualization. Set to 0 to disable.54        """55        super().__init__()56        self.backbone = backbone57        self.proposal_generator = proposal_generator58        self.roi_heads = roi_heads59 60        self.input_format = input_format61        self.vis_period = vis_period62        if vis_period > 0:63            assert input_format is not None, "input_format is required for visualization!"64 65        self.register_buffer("pixel_mean", torch.tensor(pixel_mean).view(-1, 1, 1), False)66        self.register_buffer("pixel_std", torch.tensor(pixel_std).view(-1, 1, 1), False)67        assert (68            self.pixel_mean.shape == self.pixel_std.shape69        ), f"{self.pixel_mean} and {self.pixel_std} have different shapes!"70 71    @classmethod72    def from_config(cls, cfg):73        backbone = build_backbone(cfg)74        return {75            "backbone": backbone,76            "proposal_generator": build_proposal_generator(cfg, backbone.output_shape()),77            "roi_heads": build_roi_heads(cfg, backbone.output_shape()),78            "input_format": cfg.INPUT.FORMAT,79            "vis_period": cfg.VIS_PERIOD,80            "pixel_mean": cfg.MODEL.PIXEL_MEAN,81            "pixel_std": cfg.MODEL.PIXEL_STD,82        }83 84    @property85    def device(self):86        return self.pixel_mean.device87 88    def _move_to_current_device(self, x):89        return move_device_like(x, self.pixel_mean)90 91    def visualize_training(self, batched_inputs, proposals):92        """93        A function used to visualize images and proposals. It shows ground truth94        bounding boxes on the original image and up to 20 top-scoring predicted95        object proposals on the original image. Users can implement different96        visualization functions for different models.97 98        Args:99            batched_inputs (list): a list that contains input to the model.100            proposals (list): a list that contains predicted proposals. Both101                batched_inputs and proposals should have the same length.102        """103        from detectron2.utils.visualizer import Visualizer104 105        storage = get_event_storage()106        max_vis_prop = 20107 108        for input, prop in zip(batched_inputs, proposals):109            img = input["image"]110            img = convert_image_to_rgb(img.permute(1, 2, 0), self.input_format)111            v_gt = Visualizer(img, None)112            v_gt = v_gt.overlay_instances(boxes=input["instances"].gt_boxes)113            anno_img = v_gt.get_image()114            box_size = min(len(prop.proposal_boxes), max_vis_prop)115            v_pred = Visualizer(img, None)116            v_pred = v_pred.overlay_instances(117                boxes=prop.proposal_boxes[0:box_size].tensor.cpu().numpy()118            )119            prop_img = v_pred.get_image()120            vis_img = np.concatenate((anno_img, prop_img), axis=1)121            vis_img = vis_img.transpose(2, 0, 1)122            vis_name = "Left: GT bounding boxes;  Right: Predicted proposals"123            storage.put_image(vis_name, vis_img)124            break  # only visualize one image in a batch125 126    def forward(self, batched_inputs: List[Dict[str, torch.Tensor]]):127        """128        Args:129            batched_inputs: a list, batched outputs of :class:`DatasetMapper` .130                Each item in the list contains the inputs for one image.131                For now, each item in the list is a dict that contains:132 133                * image: Tensor, image in (C, H, W) format.134                * instances (optional): groundtruth :class:`Instances`135                * proposals (optional): :class:`Instances`, precomputed proposals.136 137                Other information that's included in the original dicts, such as:138 139                * "height", "width" (int): the output resolution of the model, used in inference.140                  See :meth:`postprocess` for details.141 142        Returns:143            list[dict]:144                Each dict is the output for one input image.145                The dict contains one key "instances" whose value is a :class:`Instances`.146                The :class:`Instances` object has the following keys:147                "pred_boxes", "pred_classes", "scores", "pred_masks", "pred_keypoints"148        """149        if not self.training:150            return self.inference(batched_inputs)151 152        images = self.preprocess_image(batched_inputs)153        if "instances" in batched_inputs[0]:154            gt_instances = [x["instances"].to(self.device) for x in batched_inputs]155        else:156            gt_instances = None157 158        features = self.backbone(images.tensor)159 160        if self.proposal_generator is not None:161            proposals, proposal_losses = self.proposal_generator(images, features, gt_instances)162        else:163            assert "proposals" in batched_inputs[0]164            proposals = [x["proposals"].to(self.device) for x in batched_inputs]165            proposal_losses = {}166 167        _, detector_losses = self.roi_heads(images, features, proposals, gt_instances)168        if self.vis_period > 0:169            storage = get_event_storage()170            if storage.iter % self.vis_period == 0:171                self.visualize_training(batched_inputs, proposals)172 173        losses = {}174        losses.update(detector_losses)175        losses.update(proposal_losses)176        return losses177 178    def inference(179        self,180        batched_inputs: List[Dict[str, torch.Tensor]],181        detected_instances: Optional[List[Instances]] = None,182        do_postprocess: bool = True,183    ):184        """185        Run inference on the given inputs.186 187        Args:188            batched_inputs (list[dict]): same as in :meth:`forward`189            detected_instances (None or list[Instances]): if not None, it190                contains an `Instances` object per image. The `Instances`191                object contains "pred_boxes" and "pred_classes" which are192                known boxes in the image.193                The inference will then skip the detection of bounding boxes,194                and only predict other per-ROI outputs.195            do_postprocess (bool): whether to apply post-processing on the outputs.196 197        Returns:198            When do_postprocess=True, same as in :meth:`forward`.199            Otherwise, a list[Instances] containing raw network outputs.200        """201        assert not self.training202 203        images = self.preprocess_image(batched_inputs)204        features = self.backbone(images.tensor)205 206        if detected_instances is None:207            if self.proposal_generator is not None:208                proposals, _ = self.proposal_generator(images, features, None)209            else:210                assert "proposals" in batched_inputs[0]211                proposals = [x["proposals"].to(self.device) for x in batched_inputs]212 213            results, _ = self.roi_heads(images, features, proposals, None)214        else:215            detected_instances = [x.to(self.device) for x in detected_instances]216            results = self.roi_heads.forward_with_given_boxes(features, detected_instances)217 218        if do_postprocess:219            assert not torch.jit.is_scripting(), "Scripting is not supported for postprocess."220            return GeneralizedRCNN._postprocess(results, batched_inputs, images.image_sizes)221        return results222 223    def preprocess_image(self, batched_inputs: List[Dict[str, torch.Tensor]]):224        """225        Normalize, pad and batch the input images.226        """227        images = [self._move_to_current_device(x["image"]) for x in batched_inputs]228        images = [(x - self.pixel_mean) / self.pixel_std for x in images]229        images = ImageList.from_tensors(230            images,231            self.backbone.size_divisibility,232            padding_constraints=self.backbone.padding_constraints,233        )234        return images235 236    @staticmethod237    def _postprocess(instances, batched_inputs: List[Dict[str, torch.Tensor]], image_sizes):238        """239        Rescale the output instances to the target size.240        """241        # note: private function; subject to changes242        processed_results = []243        for results_per_image, input_per_image, image_size in zip(244            instances, batched_inputs, image_sizes245        ):246            height = input_per_image.get("height", image_size[0])247            width = input_per_image.get("width", image_size[1])248            r = detector_postprocess(results_per_image, height, width)249            processed_results.append({"instances": r})250        return processed_results251 252 253@META_ARCH_REGISTRY.register()254class ProposalNetwork(nn.Module):255    """256    A meta architecture that only predicts object proposals.257    """258 259    @configurable260    def __init__(261        self,262        *,263        backbone: Backbone,264        proposal_generator: nn.Module,265        pixel_mean: Tuple[float],266        pixel_std: Tuple[float],267    ):268        """269        Args:270            backbone: a backbone module, must follow detectron2's backbone interface271            proposal_generator: a module that generates proposals using backbone features272            pixel_mean, pixel_std: list or tuple with #channels element, representing273                the per-channel mean and std to be used to normalize the input image274        """275        super().__init__()276        self.backbone = backbone277        self.proposal_generator = proposal_generator278        self.register_buffer("pixel_mean", torch.tensor(pixel_mean).view(-1, 1, 1), False)279        self.register_buffer("pixel_std", torch.tensor(pixel_std).view(-1, 1, 1), False)280 281    @classmethod282    def from_config(cls, cfg):283        backbone = build_backbone(cfg)284        return {285            "backbone": backbone,286            "proposal_generator": build_proposal_generator(cfg, backbone.output_shape()),287            "pixel_mean": cfg.MODEL.PIXEL_MEAN,288            "pixel_std": cfg.MODEL.PIXEL_STD,289        }290 291    @property292    def device(self):293        return self.pixel_mean.device294 295    def _move_to_current_device(self, x):296        return move_device_like(x, self.pixel_mean)297 298    def forward(self, batched_inputs):299        """300        Args:301            Same as in :class:`GeneralizedRCNN.forward`302 303        Returns:304            list[dict]:305                Each dict is the output for one input image.306                The dict contains one key "proposals" whose value is a307                :class:`Instances` with keys "proposal_boxes" and "objectness_logits".308        """309        images = [self._move_to_current_device(x["image"]) for x in batched_inputs]310        images = [(x - self.pixel_mean) / self.pixel_std for x in images]311        images = ImageList.from_tensors(312            images,313            self.backbone.size_divisibility,314            padding_constraints=self.backbone.padding_constraints,315        )316        features = self.backbone(images.tensor)317 318        if "instances" in batched_inputs[0]:319            gt_instances = [x["instances"].to(self.device) for x in batched_inputs]320        elif "targets" in batched_inputs[0]:321            log_first_n(322                logging.WARN, "'targets' in the model inputs is now renamed to 'instances'!", n=10323            )324            gt_instances = [x["targets"].to(self.device) for x in batched_inputs]325        else:326            gt_instances = None327        proposals, proposal_losses = self.proposal_generator(images, features, gt_instances)328        # In training, the proposals are not useful at all but we generate them anyway.329        # This makes RPN-only models about 5% slower.330        if self.training:331            return proposal_losses332 333        processed_results = []334        for results_per_image, input_per_image, image_size in zip(335            proposals, batched_inputs, images.image_sizes336        ):337            height = input_per_image.get("height", image_size[0])338            width = input_per_image.get("width", image_size[1])339            r = detector_postprocess(results_per_image, height, width)340            processed_results.append({"proposals": r})341        return processed_results342