Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
caffe2_patch.py190 linesDownload Raw Back to export
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