Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
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 