Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import math4from typing import Dict5import torch6import torch.nn.functional as F7 8from detectron2.layers import ShapeSpec, cat9from detectron2.layers.roi_align_rotated import ROIAlignRotated10from detectron2.modeling import poolers11from detectron2.modeling.proposal_generator import rpn12from detectron2.modeling.roi_heads.mask_head import mask_rcnn_inference13from detectron2.structures import Boxes, ImageList, Instances, Keypoints, RotatedBoxes14 15from .shared import alias, to_device16 17 18"""19This file contains caffe2-compatible implementation of several detectron2 components.20"""21 22 23class Caffe2Boxes(Boxes):24 """25 Representing a list of detectron2.structures.Boxes from minibatch, each box26 is represented by a 5d vector (batch index + 4 coordinates), or a 6d vector27 (batch index + 5 coordinates) for RotatedBoxes.28 """29 30 def __init__(self, tensor):31 assert isinstance(tensor, torch.Tensor)32 assert tensor.dim() == 2 and tensor.size(-1) in [4, 5, 6], tensor.size()33 # TODO: make tensor immutable when dim is Nx5 for Boxes,34 # and Nx6 for RotatedBoxes?35 self.tensor = tensor36 37 38# TODO clean up this class, maybe just extend Instances39class InstancesList(object):40 """41 Tensor representation of a list of Instances object for a batch of images.42 43 When dealing with a batch of images with Caffe2 ops, a list of bboxes44 (instances) are usually represented by single Tensor with size45 (sigma(Ni), 5) or (sigma(Ni), 4) plus a batch split Tensor. This class is46 for providing common functions to convert between these two representations.47 """48 49 def __init__(self, im_info, indices, extra_fields=None):50 # [N, 3] -> (H, W, Scale)51 self.im_info = im_info52 # [N,] -> indice of batch to which the instance belongs53 self.indices = indices54 # [N, ...]55 self.batch_extra_fields = extra_fields or {}56 57 self.image_size = self.im_info58 59 def get_fields(self):60 """like `get_fields` in the Instances object,61 but return each field in tensor representations"""62 ret = {}63 for k, v in self.batch_extra_fields.items():64 # if isinstance(v, torch.Tensor):65 # tensor_rep = v66 # elif isinstance(v, (Boxes, Keypoints)):67 # tensor_rep = v.tensor68 # else:69 # raise ValueError("Can't find tensor representation for: {}".format())70 ret[k] = v71 return ret72 73 def has(self, name):74 return name in self.batch_extra_fields75 76 def set(self, name, value):77 # len(tensor) is a bad practice that generates ONNX constants during tracing.78 # Although not a problem for the `assert` statement below, torch ONNX exporter79 # still raises a misleading warning as it does not this call comes from `assert`80 if isinstance(value, Boxes):81 data_len = value.tensor.shape[0]82 elif isinstance(value, torch.Tensor):83 data_len = value.shape[0]84 else:85 data_len = len(value)86 if len(self.batch_extra_fields):87 assert (88 len(self) == data_len89 ), "Adding a field of length {} to a Instances of length {}".format(data_len, len(self))90 self.batch_extra_fields[name] = value91 92 def __getattr__(self, name):93 if name not in self.batch_extra_fields:94 raise AttributeError("Cannot find field '{}' in the given Instances!".format(name))95 return self.batch_extra_fields[name]96 97 def __len__(self):98 return len(self.indices)99 100 def flatten(self):101 ret = []102 for _, v in self.batch_extra_fields.items():103 if isinstance(v, (Boxes, Keypoints)):104 ret.append(v.tensor)105 else:106 ret.append(v)107 return ret108 109 @staticmethod110 def to_d2_instances_list(instances_list):111 """112 Convert InstancesList to List[Instances]. The input `instances_list` can113 also be a List[Instances], in this case this method is a non-op.114 """115 if not isinstance(instances_list, InstancesList):116 assert all(isinstance(x, Instances) for x in instances_list)117 return instances_list118 119 ret = []120 for i, info in enumerate(instances_list.im_info):121 instances = Instances(torch.Size([int(info[0].item()), int(info[1].item())]))122 123 ids = instances_list.indices == i124 for k, v in instances_list.batch_extra_fields.items():125 if isinstance(v, torch.Tensor):126 instances.set(k, v[ids])127 continue128 elif isinstance(v, Boxes):129 instances.set(k, v[ids, -4:])130 continue131 132 target_type, tensor_source = v133 assert isinstance(tensor_source, torch.Tensor)134 assert tensor_source.shape[0] == instances_list.indices.shape[0]135 tensor_source = tensor_source[ids]136 137 if issubclass(target_type, Boxes):138 instances.set(k, Boxes(tensor_source[:, -4:]))139 elif issubclass(target_type, Keypoints):140 instances.set(k, Keypoints(tensor_source))141 elif issubclass(target_type, torch.Tensor):142 instances.set(k, tensor_source)143 else:144 raise ValueError("Can't handle targe type: {}".format(target_type))145 146 ret.append(instances)147 return ret148 149 150class Caffe2Compatible(object):151 """152 A model can inherit this class to indicate that it can be traced and deployed with caffe2.153 """154 155 def _get_tensor_mode(self):156 return self._tensor_mode157 158 def _set_tensor_mode(self, v):159 self._tensor_mode = v160 161 tensor_mode = property(_get_tensor_mode, _set_tensor_mode)162 """163 If true, the model expects C2-style tensor only inputs/outputs format.164 """165 166 167class Caffe2RPN(Caffe2Compatible, rpn.RPN):168 @classmethod169 def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec]):170 ret = super(Caffe2Compatible, cls).from_config(cfg, input_shape)171 assert tuple(cfg.MODEL.RPN.BBOX_REG_WEIGHTS) == (1.0, 1.0, 1.0, 1.0) or tuple(172 cfg.MODEL.RPN.BBOX_REG_WEIGHTS173 ) == (1.0, 1.0, 1.0, 1.0, 1.0)174 return ret175 176 def _generate_proposals(177 self, images, objectness_logits_pred, anchor_deltas_pred, gt_instances=None178 ):179 assert isinstance(images, ImageList)180 if self.tensor_mode:181 im_info = images.image_sizes182 else:183 im_info = torch.tensor([[im_sz[0], im_sz[1], 1.0] for im_sz in images.image_sizes]).to(184 images.tensor.device185 )186 assert isinstance(im_info, torch.Tensor)187 188 rpn_rois_list = []189 rpn_roi_probs_list = []190 for scores, bbox_deltas, cell_anchors_tensor, feat_stride in zip(191 objectness_logits_pred,192 anchor_deltas_pred,193 [b for (n, b) in self.anchor_generator.cell_anchors.named_buffers()],194 self.anchor_generator.strides,195 ):196 scores = scores.detach()197 bbox_deltas = bbox_deltas.detach()198 199 rpn_rois, rpn_roi_probs = torch.ops._caffe2.GenerateProposals(200 scores,201 bbox_deltas,202 im_info,203 cell_anchors_tensor,204 spatial_scale=1.0 / feat_stride,205 pre_nms_topN=self.pre_nms_topk[self.training],206 post_nms_topN=self.post_nms_topk[self.training],207 nms_thresh=self.nms_thresh,208 min_size=self.min_box_size,209 # correct_transform_coords=True, # deprecated argument210 angle_bound_on=True, # Default211 angle_bound_lo=-180,212 angle_bound_hi=180,213 clip_angle_thresh=1.0, # Default214 legacy_plus_one=False,215 )216 rpn_rois_list.append(rpn_rois)217 rpn_roi_probs_list.append(rpn_roi_probs)218 219 # For FPN in D2, in RPN all proposals from different levels are concated220 # together, ranked and picked by top post_nms_topk. Then in ROIPooler221 # it calculates level_assignments and calls the RoIAlign from222 # the corresponding level.223 224 if len(objectness_logits_pred) == 1:225 rpn_rois = rpn_rois_list[0]226 rpn_roi_probs = rpn_roi_probs_list[0]227 else:228 assert len(rpn_rois_list) == len(rpn_roi_probs_list)229 rpn_post_nms_topN = self.post_nms_topk[self.training]230 231 device = rpn_rois_list[0].device232 input_list = [to_device(x, "cpu") for x in (rpn_rois_list + rpn_roi_probs_list)]233 234 # TODO remove this after confirming rpn_max_level/rpn_min_level235 # is not needed in CollectRpnProposals.236 feature_strides = list(self.anchor_generator.strides)237 rpn_min_level = int(math.log2(feature_strides[0]))238 rpn_max_level = int(math.log2(feature_strides[-1]))239 assert (rpn_max_level - rpn_min_level + 1) == len(240 rpn_rois_list241 ), "CollectRpnProposals requires continuous levels"242 243 rpn_rois = torch.ops._caffe2.CollectRpnProposals(244 input_list,245 # NOTE: in current implementation, rpn_max_level and rpn_min_level246 # are not needed, only the subtraction of two matters and it247 # can be infer from the number of inputs. Keep them now for248 # consistency.249 rpn_max_level=2 + len(rpn_rois_list) - 1,250 rpn_min_level=2,251 rpn_post_nms_topN=rpn_post_nms_topN,252 )253 rpn_rois = to_device(rpn_rois, device)254 rpn_roi_probs = []255 256 proposals = self.c2_postprocess(im_info, rpn_rois, rpn_roi_probs, self.tensor_mode)257 return proposals, {}258 259 def forward(self, images, features, gt_instances=None):260 assert not self.training261 features = [features[f] for f in self.in_features]262 objectness_logits_pred, anchor_deltas_pred = self.rpn_head(features)263 return self._generate_proposals(264 images,265 objectness_logits_pred,266 anchor_deltas_pred,267 gt_instances,268 )269 270 @staticmethod271 def c2_postprocess(im_info, rpn_rois, rpn_roi_probs, tensor_mode):272 proposals = InstancesList(273 im_info=im_info,274 indices=rpn_rois[:, 0],275 extra_fields={276 "proposal_boxes": Caffe2Boxes(rpn_rois),277 "objectness_logits": (torch.Tensor, rpn_roi_probs),278 },279 )280 if not tensor_mode:281 proposals = InstancesList.to_d2_instances_list(proposals)282 else:283 proposals = [proposals]284 return proposals285 286 287class Caffe2ROIPooler(Caffe2Compatible, poolers.ROIPooler):288 @staticmethod289 def c2_preprocess(box_lists):290 assert all(isinstance(x, Boxes) for x in box_lists)291 if all(isinstance(x, Caffe2Boxes) for x in box_lists):292 # input is pure-tensor based293 assert len(box_lists) == 1294 pooler_fmt_boxes = box_lists[0].tensor295 else:296 pooler_fmt_boxes = poolers.convert_boxes_to_pooler_format(box_lists)297 return pooler_fmt_boxes298 299 def forward(self, x, box_lists):300 assert not self.training301 302 pooler_fmt_boxes = self.c2_preprocess(box_lists)303 num_level_assignments = len(self.level_poolers)304 305 if num_level_assignments == 1:306 if isinstance(self.level_poolers[0], ROIAlignRotated):307 c2_roi_align = torch.ops._caffe2.RoIAlignRotated308 aligned = True309 else:310 c2_roi_align = torch.ops._caffe2.RoIAlign311 aligned = self.level_poolers[0].aligned312 313 x0 = x[0]314 if x0.is_quantized:315 x0 = x0.dequantize()316 317 out = c2_roi_align(318 x0,319 pooler_fmt_boxes,320 order="NCHW",321 spatial_scale=float(self.level_poolers[0].spatial_scale),322 pooled_h=int(self.output_size[0]),323 pooled_w=int(self.output_size[1]),324 sampling_ratio=int(self.level_poolers[0].sampling_ratio),325 aligned=aligned,326 )327 return out328 329 device = pooler_fmt_boxes.device330 assert (331 self.max_level - self.min_level + 1 == 4332 ), "Currently DistributeFpnProposals only support 4 levels"333 fpn_outputs = torch.ops._caffe2.DistributeFpnProposals(334 to_device(pooler_fmt_boxes, "cpu"),335 roi_canonical_scale=self.canonical_box_size,336 roi_canonical_level=self.canonical_level,337 roi_max_level=self.max_level,338 roi_min_level=self.min_level,339 legacy_plus_one=False,340 )341 fpn_outputs = [to_device(x, device) for x in fpn_outputs]342 343 rois_fpn_list = fpn_outputs[:-1]344 rois_idx_restore_int32 = fpn_outputs[-1]345 346 roi_feat_fpn_list = []347 for roi_fpn, x_level, pooler in zip(rois_fpn_list, x, self.level_poolers):348 if isinstance(pooler, ROIAlignRotated):349 c2_roi_align = torch.ops._caffe2.RoIAlignRotated350 aligned = True351 else:352 c2_roi_align = torch.ops._caffe2.RoIAlign353 aligned = bool(pooler.aligned)354 355 if x_level.is_quantized:356 x_level = x_level.dequantize()357 358 roi_feat_fpn = c2_roi_align(359 x_level,360 roi_fpn,361 order="NCHW",362 spatial_scale=float(pooler.spatial_scale),363 pooled_h=int(self.output_size[0]),364 pooled_w=int(self.output_size[1]),365 sampling_ratio=int(pooler.sampling_ratio),366 aligned=aligned,367 )368 roi_feat_fpn_list.append(roi_feat_fpn)369 370 roi_feat_shuffled = cat(roi_feat_fpn_list, dim=0)371 assert roi_feat_shuffled.numel() > 0 and rois_idx_restore_int32.numel() > 0, (372 "Caffe2 export requires tracing with a model checkpoint + input that can produce valid"373 " detections. But no detections were obtained with the given checkpoint and input!"374 )375 roi_feat = torch.ops._caffe2.BatchPermutation(roi_feat_shuffled, rois_idx_restore_int32)376 return roi_feat377 378 379def caffe2_fast_rcnn_outputs_inference(tensor_mode, box_predictor, predictions, proposals):380 """equivalent to FastRCNNOutputLayers.inference"""381 num_classes = box_predictor.num_classes382 score_thresh = box_predictor.test_score_thresh383 nms_thresh = box_predictor.test_nms_thresh384 topk_per_image = box_predictor.test_topk_per_image385 is_rotated = len(box_predictor.box2box_transform.weights) == 5386 387 if is_rotated:388 box_dim = 5389 assert box_predictor.box2box_transform.weights[4] == 1, (390 "The weights for Rotated BBoxTransform in C2 have only 4 dimensions,"391 + " thus enforcing the angle weight to be 1 for now"392 )393 box2box_transform_weights = box_predictor.box2box_transform.weights[:4]394 else:395 box_dim = 4396 box2box_transform_weights = box_predictor.box2box_transform.weights397 398 class_logits, box_regression = predictions399 if num_classes + 1 == class_logits.shape[1]:400 class_prob = F.softmax(class_logits, -1)401 else:402 assert num_classes == class_logits.shape[1]403 class_prob = F.sigmoid(class_logits)404 # BoxWithNMSLimit will infer num_classes from the shape of the class_prob405 # So append a zero column as placeholder for the background class406 class_prob = torch.cat((class_prob, torch.zeros(class_prob.shape[0], 1)), dim=1)407 408 assert box_regression.shape[1] % box_dim == 0409 cls_agnostic_bbox_reg = box_regression.shape[1] // box_dim == 1410 411 input_tensor_mode = proposals[0].proposal_boxes.tensor.shape[1] == box_dim + 1412 413 proposal_boxes = proposals[0].proposal_boxes414 if isinstance(proposal_boxes, Caffe2Boxes):415 rois = Caffe2Boxes.cat([p.proposal_boxes for p in proposals])416 elif isinstance(proposal_boxes, RotatedBoxes):417 rois = RotatedBoxes.cat([p.proposal_boxes for p in proposals])418 elif isinstance(proposal_boxes, Boxes):419 rois = Boxes.cat([p.proposal_boxes for p in proposals])420 else:421 raise NotImplementedError(422 'Expected proposals[0].proposal_boxes to be type "Boxes", '423 f"instead got {type(proposal_boxes)}"424 )425 426 device, dtype = rois.tensor.device, rois.tensor.dtype427 if input_tensor_mode:428 im_info = proposals[0].image_size429 rois = rois.tensor430 else:431 im_info = torch.tensor([[sz[0], sz[1], 1.0] for sz in [x.image_size for x in proposals]])432 batch_ids = cat(433 [434 torch.full((b, 1), i, dtype=dtype, device=device)435 for i, b in enumerate(len(p) for p in proposals)436 ],437 dim=0,438 )439 rois = torch.cat([batch_ids, rois.tensor], dim=1)440 441 roi_pred_bbox, roi_batch_splits = torch.ops._caffe2.BBoxTransform(442 to_device(rois, "cpu"),443 to_device(box_regression, "cpu"),444 to_device(im_info, "cpu"),445 weights=box2box_transform_weights,446 apply_scale=True,447 rotated=is_rotated,448 angle_bound_on=True,449 angle_bound_lo=-180,450 angle_bound_hi=180,451 clip_angle_thresh=1.0,452 legacy_plus_one=False,453 )454 roi_pred_bbox = to_device(roi_pred_bbox, device)455 roi_batch_splits = to_device(roi_batch_splits, device)456 457 nms_outputs = torch.ops._caffe2.BoxWithNMSLimit(458 to_device(class_prob, "cpu"),459 to_device(roi_pred_bbox, "cpu"),460 to_device(roi_batch_splits, "cpu"),461 score_thresh=float(score_thresh),462 nms=float(nms_thresh),463 detections_per_im=int(topk_per_image),464 soft_nms_enabled=False,465 soft_nms_method="linear",466 soft_nms_sigma=0.5,467 soft_nms_min_score_thres=0.001,468 rotated=is_rotated,469 cls_agnostic_bbox_reg=cls_agnostic_bbox_reg,470 input_boxes_include_bg_cls=False,471 output_classes_include_bg_cls=False,472 legacy_plus_one=False,473 )474 roi_score_nms = to_device(nms_outputs[0], device)475 roi_bbox_nms = to_device(nms_outputs[1], device)476 roi_class_nms = to_device(nms_outputs[2], device)477 roi_batch_splits_nms = to_device(nms_outputs[3], device)478 roi_keeps_nms = to_device(nms_outputs[4], device)479 roi_keeps_size_nms = to_device(nms_outputs[5], device)480 if not tensor_mode:481 roi_class_nms = roi_class_nms.to(torch.int64)482 483 roi_batch_ids = cat(484 [485 torch.full((b, 1), i, dtype=dtype, device=device)486 for i, b in enumerate(int(x.item()) for x in roi_batch_splits_nms)487 ],488 dim=0,489 )490 491 roi_class_nms = alias(roi_class_nms, "class_nms")492 roi_score_nms = alias(roi_score_nms, "score_nms")493 roi_bbox_nms = alias(roi_bbox_nms, "bbox_nms")494 roi_batch_splits_nms = alias(roi_batch_splits_nms, "batch_splits_nms")495 roi_keeps_nms = alias(roi_keeps_nms, "keeps_nms")496 roi_keeps_size_nms = alias(roi_keeps_size_nms, "keeps_size_nms")497 498 results = InstancesList(499 im_info=im_info,500 indices=roi_batch_ids[:, 0],501 extra_fields={502 "pred_boxes": Caffe2Boxes(roi_bbox_nms),503 "scores": roi_score_nms,504 "pred_classes": roi_class_nms,505 },506 )507 508 if not tensor_mode:509 results = InstancesList.to_d2_instances_list(results)510 batch_splits = roi_batch_splits_nms.int().tolist()511 kept_indices = list(roi_keeps_nms.to(torch.int64).split(batch_splits))512 else:513 results = [results]514 kept_indices = [roi_keeps_nms]515 516 return results, kept_indices517 518 519class Caffe2FastRCNNOutputsInference:520 def __init__(self, tensor_mode):521 self.tensor_mode = tensor_mode # whether the output is caffe2 tensor mode522 523 def __call__(self, box_predictor, predictions, proposals):524 return caffe2_fast_rcnn_outputs_inference(525 self.tensor_mode, box_predictor, predictions, proposals526 )527 528 529def caffe2_mask_rcnn_inference(pred_mask_logits, pred_instances):530 """equivalent to mask_head.mask_rcnn_inference"""531 if all(isinstance(x, InstancesList) for x in pred_instances):532 assert len(pred_instances) == 1533 mask_probs_pred = pred_mask_logits.sigmoid()534 mask_probs_pred = alias(mask_probs_pred, "mask_fcn_probs")535 pred_instances[0].set("pred_masks", mask_probs_pred)536 else:537 mask_rcnn_inference(pred_mask_logits, pred_instances)538 539 540class Caffe2MaskRCNNInference:541 def __call__(self, pred_mask_logits, pred_instances):542 return caffe2_mask_rcnn_inference(pred_mask_logits, pred_instances)543 544 545def caffe2_keypoint_rcnn_inference(use_heatmap_max_keypoint, pred_keypoint_logits, pred_instances):546 # just return the keypoint heatmap for now,547 # there will be option to call HeatmapMaxKeypointOp548 output = alias(pred_keypoint_logits, "kps_score")549 if all(isinstance(x, InstancesList) for x in pred_instances):550 assert len(pred_instances) == 1551 if use_heatmap_max_keypoint:552 device = output.device553 output = torch.ops._caffe2.HeatmapMaxKeypoint(554 to_device(output, "cpu"),555 pred_instances[0].pred_boxes.tensor,556 should_output_softmax=True, # worth make it configerable?557 )558 output = to_device(output, device)559 output = alias(output, "keypoints_out")560 pred_instances[0].set("pred_keypoints", output)561 return pred_keypoint_logits562 563 564class Caffe2KeypointRCNNInference:565 def __init__(self, use_heatmap_max_keypoint):566 self.use_heatmap_max_keypoint = use_heatmap_max_keypoint567 568 def __call__(self, pred_keypoint_logits, pred_instances):569 return caffe2_keypoint_rcnn_inference(570 self.use_heatmap_max_keypoint, pred_keypoint_logits, pred_instances571 )572 