Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import contextlib4from unittest import mock5import torch6 7from detectron2.modeling import poolers8from detectron2.modeling.proposal_generator import rpn9from detectron2.modeling.roi_heads import keypoint_head, mask_head10from detectron2.modeling.roi_heads.fast_rcnn import FastRCNNOutputLayers11 12from .c10 import (13 Caffe2Compatible,14 Caffe2FastRCNNOutputsInference,15 Caffe2KeypointRCNNInference,16 Caffe2MaskRCNNInference,17 Caffe2ROIPooler,18 Caffe2RPN,19 caffe2_fast_rcnn_outputs_inference,20 caffe2_keypoint_rcnn_inference,21 caffe2_mask_rcnn_inference,22)23 24 25class GenericMixin(object):26 pass27 28 29class Caffe2CompatibleConverter(object):30 """31 A GenericUpdater which implements the `create_from` interface, by modifying32 module object and assign it with another class replaceCls.33 """34 35 def __init__(self, replaceCls):36 self.replaceCls = replaceCls37 38 def create_from(self, module):39 # update module's class to the new class40 assert isinstance(module, torch.nn.Module)41 if issubclass(self.replaceCls, GenericMixin):42 # replaceCls should act as mixin, create a new class on-the-fly43 new_class = type(44 "{}MixedWith{}".format(self.replaceCls.__name__, module.__class__.__name__),45 (self.replaceCls, module.__class__),46 {}, # {"new_method": lambda self: ...},47 )48 module.__class__ = new_class49 else:50 # replaceCls is complete class, this allow arbitrary class swap51 module.__class__ = self.replaceCls52 53 # initialize Caffe2Compatible54 if isinstance(module, Caffe2Compatible):55 module.tensor_mode = False56 57 return module58 59 60def patch(model, target, updater, *args, **kwargs):61 """62 recursively (post-order) update all modules with the target type and its63 subclasses, make a initialization/composition/inheritance/... via the64 updater.create_from.65 """66 for name, module in model.named_children():67 model._modules[name] = patch(module, target, updater, *args, **kwargs)68 if isinstance(model, target):69 return updater.create_from(model, *args, **kwargs)70 return model71 72 73def patch_generalized_rcnn(model):74 ccc = Caffe2CompatibleConverter75 model = patch(model, rpn.RPN, ccc(Caffe2RPN))76 model = patch(model, poolers.ROIPooler, ccc(Caffe2ROIPooler))77 78 return model79 80 81@contextlib.contextmanager82def mock_fastrcnn_outputs_inference(83 tensor_mode, check=True, box_predictor_type=FastRCNNOutputLayers84):85 with mock.patch.object(86 box_predictor_type,87 "inference",88 autospec=True,89 side_effect=Caffe2FastRCNNOutputsInference(tensor_mode),90 ) as mocked_func:91 yield92 if check:93 assert mocked_func.call_count > 094 95 96@contextlib.contextmanager97def mock_mask_rcnn_inference(tensor_mode, patched_module, check=True):98 with mock.patch(99 "{}.mask_rcnn_inference".format(patched_module), side_effect=Caffe2MaskRCNNInference()100 ) as mocked_func:101 yield102 if check:103 assert mocked_func.call_count > 0104 105 106@contextlib.contextmanager107def mock_keypoint_rcnn_inference(tensor_mode, patched_module, use_heatmap_max_keypoint, check=True):108 with mock.patch(109 "{}.keypoint_rcnn_inference".format(patched_module),110 side_effect=Caffe2KeypointRCNNInference(use_heatmap_max_keypoint),111 ) as mocked_func:112 yield113 if check:114 assert mocked_func.call_count > 0115 116 117class ROIHeadsPatcher:118 def __init__(self, heads, use_heatmap_max_keypoint):119 self.heads = heads120 self.use_heatmap_max_keypoint = use_heatmap_max_keypoint121 self.previous_patched = {}122 123 @contextlib.contextmanager124 def mock_roi_heads(self, tensor_mode=True):125 """126 Patching several inference functions inside ROIHeads and its subclasses127 128 Args:129 tensor_mode (bool): whether the inputs/outputs are caffe2's tensor130 format or not. Default to True.131 """132 # NOTE: this requries the `keypoint_rcnn_inference` and `mask_rcnn_inference`133 # are called inside the same file as BaseXxxHead due to using mock.patch.134 kpt_heads_mod = keypoint_head.BaseKeypointRCNNHead.__module__135 mask_head_mod = mask_head.BaseMaskRCNNHead.__module__136 137 mock_ctx_managers = [138 mock_fastrcnn_outputs_inference(139 tensor_mode=tensor_mode,140 check=True,141 box_predictor_type=type(self.heads.box_predictor),142 )143 ]144 if getattr(self.heads, "keypoint_on", False):145 mock_ctx_managers += [146 mock_keypoint_rcnn_inference(147 tensor_mode, kpt_heads_mod, self.use_heatmap_max_keypoint148 )149 ]150 if getattr(self.heads, "mask_on", False):151 mock_ctx_managers += [mock_mask_rcnn_inference(tensor_mode, mask_head_mod)]152 153 with contextlib.ExitStack() as stack: # python 3.3+154 for mgr in mock_ctx_managers:155 stack.enter_context(mgr)156 yield157 158 def patch_roi_heads(self, tensor_mode=True):159 self.previous_patched["box_predictor"] = self.heads.box_predictor.inference160 self.previous_patched["keypoint_rcnn"] = keypoint_head.keypoint_rcnn_inference161 self.previous_patched["mask_rcnn"] = mask_head.mask_rcnn_inference162 163 def patched_fastrcnn_outputs_inference(predictions, proposal):164 return caffe2_fast_rcnn_outputs_inference(165 True, self.heads.box_predictor, predictions, proposal166 )167 168 self.heads.box_predictor.inference = patched_fastrcnn_outputs_inference169 170 if getattr(self.heads, "keypoint_on", False):171 172 def patched_keypoint_rcnn_inference(pred_keypoint_logits, pred_instances):173 return caffe2_keypoint_rcnn_inference(174 self.use_heatmap_max_keypoint, pred_keypoint_logits, pred_instances175 )176 177 keypoint_head.keypoint_rcnn_inference = patched_keypoint_rcnn_inference178 179 if getattr(self.heads, "mask_on", False):180 181 def patched_mask_rcnn_inference(pred_mask_logits, pred_instances):182 return caffe2_mask_rcnn_inference(pred_mask_logits, pred_instances)183 184 mask_head.mask_rcnn_inference = patched_mask_rcnn_inference185 186 def unpatch_roi_heads(self):187 self.heads.box_predictor.inference = self.previous_patched["box_predictor"]188 keypoint_head.keypoint_rcnn_inference = self.previous_patched["keypoint_rcnn"]189 mask_head.mask_rcnn_inference = self.previous_patched["mask_rcnn"]190 