Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
caffe2_modeling.py421 linesDownload Raw Back to export
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import functools4import io5import struct6import types7import torch8 9from detectron2.modeling import meta_arch10from detectron2.modeling.box_regression import Box2BoxTransform11from detectron2.modeling.roi_heads import keypoint_head12from detectron2.structures import Boxes, ImageList, Instances, RotatedBoxes13 14from .c10 import Caffe2Compatible15from .caffe2_patch import ROIHeadsPatcher, patch_generalized_rcnn16from .shared import (17    alias,18    check_set_pb_arg,19    get_pb_arg_floats,20    get_pb_arg_valf,21    get_pb_arg_vali,22    get_pb_arg_vals,23    mock_torch_nn_functional_interpolate,24)25 26 27def assemble_rcnn_outputs_by_name(image_sizes, tensor_outputs, force_mask_on=False):28    """29    A function to assemble caffe2 model's outputs (i.e. Dict[str, Tensor])30    to detectron2's format (i.e. list of Instances instance).31    This only works when the model follows the Caffe2 detectron's naming convention.32 33    Args:34        image_sizes (List[List[int, int]]): [H, W] of every image.35        tensor_outputs (Dict[str, Tensor]): external_output to its tensor.36 37        force_mask_on (Bool): if true, the it make sure there'll be pred_masks even38            if the mask is not found from tensor_outputs (usually due to model crash)39    """40 41    results = [Instances(image_size) for image_size in image_sizes]42 43    batch_splits = tensor_outputs.get("batch_splits", None)44    if batch_splits:45        raise NotImplementedError()46    assert len(image_sizes) == 147    result = results[0]48 49    bbox_nms = tensor_outputs["bbox_nms"]50    score_nms = tensor_outputs["score_nms"]51    class_nms = tensor_outputs["class_nms"]52    # Detection will always success because Conv support 0-batch53    assert bbox_nms is not None54    assert score_nms is not None55    assert class_nms is not None56    if bbox_nms.shape[1] == 5:57        result.pred_boxes = RotatedBoxes(bbox_nms)58    else:59        result.pred_boxes = Boxes(bbox_nms)60    result.scores = score_nms61    result.pred_classes = class_nms.to(torch.int64)62 63    mask_fcn_probs = tensor_outputs.get("mask_fcn_probs", None)64    if mask_fcn_probs is not None:65        # finish the mask pred66        mask_probs_pred = mask_fcn_probs67        num_masks = mask_probs_pred.shape[0]68        class_pred = result.pred_classes69        indices = torch.arange(num_masks, device=class_pred.device)70        mask_probs_pred = mask_probs_pred[indices, class_pred][:, None]71        result.pred_masks = mask_probs_pred72    elif force_mask_on:73        # NOTE: there's no way to know the height/width of mask here, it won't be74        # used anyway when batch size is 0, so just set them to 0.75        result.pred_masks = torch.zeros([0, 1, 0, 0], dtype=torch.uint8)76 77    keypoints_out = tensor_outputs.get("keypoints_out", None)78    kps_score = tensor_outputs.get("kps_score", None)79    if keypoints_out is not None:80        # keypoints_out: [N, 4, #kypoints], where 4 is in order of (x, y, score, prob)81        keypoints_tensor = keypoints_out82        # NOTE: it's possible that prob is not calculated if "should_output_softmax"83        # is set to False in HeatmapMaxKeypoint, so just using raw score, seems84        # it doesn't affect mAP. TODO: check more carefully.85        keypoint_xyp = keypoints_tensor.transpose(1, 2)[:, :, [0, 1, 2]]86        result.pred_keypoints = keypoint_xyp87    elif kps_score is not None:88        # keypoint heatmap to sparse data structure89        pred_keypoint_logits = kps_score90        keypoint_head.keypoint_rcnn_inference(pred_keypoint_logits, [result])91 92    return results93 94 95def _cast_to_f32(f64):96    return struct.unpack("f", struct.pack("f", f64))[0]97 98 99def set_caffe2_compatible_tensor_mode(model, enable=True):100    def _fn(m):101        if isinstance(m, Caffe2Compatible):102            m.tensor_mode = enable103 104    model.apply(_fn)105 106 107def convert_batched_inputs_to_c2_format(batched_inputs, size_divisibility, device):108    """109    See get_caffe2_inputs() below.110    """111    assert all(isinstance(x, dict) for x in batched_inputs)112    assert all(x["image"].dim() == 3 for x in batched_inputs)113 114    images = [x["image"] for x in batched_inputs]115    images = ImageList.from_tensors(images, size_divisibility)116 117    im_info = []118    for input_per_image, image_size in zip(batched_inputs, images.image_sizes):119        target_height = input_per_image.get("height", image_size[0])120        target_width = input_per_image.get("width", image_size[1])  # noqa121        # NOTE: The scale inside im_info is kept as convention and for providing122        # post-processing information if further processing is needed. For123        # current Caffe2 model definitions that don't include post-processing inside124        # the model, this number is not used.125        # NOTE: There can be a slight difference between width and height126        # scales, using a single number can results in numerical difference127        # compared with D2's post-processing.128        scale = target_height / image_size[0]129        im_info.append([image_size[0], image_size[1], scale])130    im_info = torch.Tensor(im_info)131 132    return images.tensor.to(device), im_info.to(device)133 134 135class Caffe2MetaArch(Caffe2Compatible, torch.nn.Module):136    """137    Base class for caffe2-compatible implementation of a meta architecture.138    The forward is traceable and its traced graph can be converted to caffe2139    graph through ONNX.140    """141 142    def __init__(self, cfg, torch_model, enable_tensor_mode=True):143        """144        Args:145            cfg (CfgNode):146            torch_model (nn.Module): the detectron2 model (meta_arch) to be147                converted.148        """149        super().__init__()150        self._wrapped_model = torch_model151        self.eval()152        set_caffe2_compatible_tensor_mode(self, enable_tensor_mode)153 154    def get_caffe2_inputs(self, batched_inputs):155        """156        Convert pytorch-style structured inputs to caffe2-style inputs that157        are tuples of tensors.158 159        Args:160            batched_inputs (list[dict]): inputs to a detectron2 model161                in its standard format. Each dict has "image" (CHW tensor), and optionally162                "height" and "width".163 164        Returns:165            tuple[Tensor]:166                tuple of tensors that will be the inputs to the167                :meth:`forward` method. For existing models, the first168                is an NCHW tensor (padded and batched); the second is169                a im_info Nx3 tensor, where the rows are170                (height, width, unused legacy parameter)171        """172        return convert_batched_inputs_to_c2_format(173            batched_inputs,174            self._wrapped_model.backbone.size_divisibility,175            self._wrapped_model.device,176        )177 178    def encode_additional_info(self, predict_net, init_net):179        """180        Save extra metadata that will be used by inference in the output protobuf.181        """182        pass183 184    def forward(self, inputs):185        """186        Run the forward in caffe2-style. It has to use caffe2-compatible ops187        and the method will be used for tracing.188 189        Args:190            inputs (tuple[Tensor]): inputs defined by :meth:`get_caffe2_input`.191                They will be the inputs of the converted caffe2 graph.192 193        Returns:194            tuple[Tensor]: output tensors. They will be the outputs of the195                converted caffe2 graph.196        """197        raise NotImplementedError198 199    def _caffe2_preprocess_image(self, inputs):200        """201        Caffe2 implementation of preprocess_image, which is called inside each MetaArch's forward.202        It normalizes the input images, and the final caffe2 graph assumes the203        inputs have been batched already.204        """205        data, im_info = inputs206        data = alias(data, "data")207        im_info = alias(im_info, "im_info")208        mean, std = self._wrapped_model.pixel_mean, self._wrapped_model.pixel_std209        normalized_data = (data - mean) / std210        normalized_data = alias(normalized_data, "normalized_data")211 212        # Pack (data, im_info) into ImageList which is recognized by self.inference.213        images = ImageList(tensor=normalized_data, image_sizes=im_info)214        return images215 216    @staticmethod217    def get_outputs_converter(predict_net, init_net):218        """219        Creates a function that converts outputs of the caffe2 model to220        detectron2's standard format.221        The function uses information in `predict_net` and `init_net` that are222        available at inferene time. Therefore the function logic can be used in inference.223 224        The returned function has the following signature:225 226            def convert(batched_inputs, c2_inputs, c2_results) -> detectron2_outputs227 228        Where229 230            * batched_inputs (list[dict]): the original input format of the meta arch231            * c2_inputs (tuple[Tensor]): the caffe2 inputs.232            * c2_results (dict[str, Tensor]): the caffe2 output format,233                corresponding to the outputs of the :meth:`forward` function.234            * detectron2_outputs: the original output format of the meta arch.235 236        This function can be used to compare the outputs of the original meta arch and237        the converted caffe2 graph.238 239        Returns:240            callable: a callable of the above signature.241        """242        raise NotImplementedError243 244 245class Caffe2GeneralizedRCNN(Caffe2MetaArch):246    def __init__(self, cfg, torch_model, enable_tensor_mode=True):247        assert isinstance(torch_model, meta_arch.GeneralizedRCNN)248        torch_model = patch_generalized_rcnn(torch_model)249        super().__init__(cfg, torch_model, enable_tensor_mode)250 251        try:252            use_heatmap_max_keypoint = cfg.EXPORT_CAFFE2.USE_HEATMAP_MAX_KEYPOINT253        except AttributeError:254            use_heatmap_max_keypoint = False255        self.roi_heads_patcher = ROIHeadsPatcher(256            self._wrapped_model.roi_heads, use_heatmap_max_keypoint257        )258        if self.tensor_mode:259            self.roi_heads_patcher.patch_roi_heads()260 261    def encode_additional_info(self, predict_net, init_net):262        size_divisibility = self._wrapped_model.backbone.size_divisibility263        check_set_pb_arg(predict_net, "size_divisibility", "i", size_divisibility)264        check_set_pb_arg(265            predict_net, "device", "s", str.encode(str(self._wrapped_model.device), "ascii")266        )267        check_set_pb_arg(predict_net, "meta_architecture", "s", b"GeneralizedRCNN")268 269    @mock_torch_nn_functional_interpolate()270    def forward(self, inputs):271        if not self.tensor_mode:272            return self._wrapped_model.inference(inputs)273        images = self._caffe2_preprocess_image(inputs)274        features = self._wrapped_model.backbone(images.tensor)275        proposals, _ = self._wrapped_model.proposal_generator(images, features)276        detector_results, _ = self._wrapped_model.roi_heads(images, features, proposals)277        return tuple(detector_results[0].flatten())278 279    @staticmethod280    def get_outputs_converter(predict_net, init_net):281        def f(batched_inputs, c2_inputs, c2_results):282            _, im_info = c2_inputs283            image_sizes = [[int(im[0]), int(im[1])] for im in im_info]284            results = assemble_rcnn_outputs_by_name(image_sizes, c2_results)285            return meta_arch.GeneralizedRCNN._postprocess(results, batched_inputs, image_sizes)286 287        return f288 289 290class Caffe2RetinaNet(Caffe2MetaArch):291    def __init__(self, cfg, torch_model):292        assert isinstance(torch_model, meta_arch.RetinaNet)293        super().__init__(cfg, torch_model)294 295    @mock_torch_nn_functional_interpolate()296    def forward(self, inputs):297        assert self.tensor_mode298        images = self._caffe2_preprocess_image(inputs)299 300        # explicitly return the images sizes to avoid removing "im_info" by ONNX301        # since it's not used in the forward path302        return_tensors = [images.image_sizes]303 304        features = self._wrapped_model.backbone(images.tensor)305        features = [features[f] for f in self._wrapped_model.head_in_features]306        for i, feature_i in enumerate(features):307            features[i] = alias(feature_i, "feature_{}".format(i), is_backward=True)308            return_tensors.append(features[i])309 310        pred_logits, pred_anchor_deltas = self._wrapped_model.head(features)311        for i, (box_cls_i, box_delta_i) in enumerate(zip(pred_logits, pred_anchor_deltas)):312            return_tensors.append(alias(box_cls_i, "box_cls_{}".format(i)))313            return_tensors.append(alias(box_delta_i, "box_delta_{}".format(i)))314 315        return tuple(return_tensors)316 317    def encode_additional_info(self, predict_net, init_net):318        size_divisibility = self._wrapped_model.backbone.size_divisibility319        check_set_pb_arg(predict_net, "size_divisibility", "i", size_divisibility)320        check_set_pb_arg(321            predict_net, "device", "s", str.encode(str(self._wrapped_model.device), "ascii")322        )323        check_set_pb_arg(predict_net, "meta_architecture", "s", b"RetinaNet")324 325        # Inference parameters:326        check_set_pb_arg(327            predict_net, "score_threshold", "f", _cast_to_f32(self._wrapped_model.test_score_thresh)328        )329        check_set_pb_arg(330            predict_net, "topk_candidates", "i", self._wrapped_model.test_topk_candidates331        )332        check_set_pb_arg(333            predict_net, "nms_threshold", "f", _cast_to_f32(self._wrapped_model.test_nms_thresh)334        )335        check_set_pb_arg(336            predict_net,337            "max_detections_per_image",338            "i",339            self._wrapped_model.max_detections_per_image,340        )341 342        check_set_pb_arg(343            predict_net,344            "bbox_reg_weights",345            "floats",346            [_cast_to_f32(w) for w in self._wrapped_model.box2box_transform.weights],347        )348        self._encode_anchor_generator_cfg(predict_net)349 350    def _encode_anchor_generator_cfg(self, predict_net):351        # serialize anchor_generator for future use352        serialized_anchor_generator = io.BytesIO()353        torch.save(self._wrapped_model.anchor_generator, serialized_anchor_generator)354        # Ideally we can put anchor generating inside the model, then we don't355        # need to store this information.356        bytes = serialized_anchor_generator.getvalue()357        check_set_pb_arg(predict_net, "serialized_anchor_generator", "s", bytes)358 359    @staticmethod360    def get_outputs_converter(predict_net, init_net):361        self = types.SimpleNamespace()362        serialized_anchor_generator = io.BytesIO(363            get_pb_arg_vals(predict_net, "serialized_anchor_generator", None)364        )365        self.anchor_generator = torch.load(serialized_anchor_generator)366        bbox_reg_weights = get_pb_arg_floats(predict_net, "bbox_reg_weights", None)367        self.box2box_transform = Box2BoxTransform(weights=tuple(bbox_reg_weights))368        self.test_score_thresh = get_pb_arg_valf(predict_net, "score_threshold", None)369        self.test_topk_candidates = get_pb_arg_vali(predict_net, "topk_candidates", None)370        self.test_nms_thresh = get_pb_arg_valf(predict_net, "nms_threshold", None)371        self.max_detections_per_image = get_pb_arg_vali(372            predict_net, "max_detections_per_image", None373        )374 375        # hack to reuse inference code from RetinaNet376        for meth in [377            "forward_inference",378            "inference_single_image",379            "_transpose_dense_predictions",380            "_decode_multi_level_predictions",381            "_decode_per_level_predictions",382        ]:383            setattr(self, meth, functools.partial(getattr(meta_arch.RetinaNet, meth), self))384 385        def f(batched_inputs, c2_inputs, c2_results):386            _, im_info = c2_inputs387            image_sizes = [[int(im[0]), int(im[1])] for im in im_info]388            dummy_images = ImageList(389                torch.randn(390                    (391                        len(im_info),392                        3,393                    )394                    + tuple(image_sizes[0])395                ),396                image_sizes,397            )398 399            num_features = len([x for x in c2_results.keys() if x.startswith("box_cls_")])400            pred_logits = [c2_results["box_cls_{}".format(i)] for i in range(num_features)]401            pred_anchor_deltas = [c2_results["box_delta_{}".format(i)] for i in range(num_features)]402 403            # For each feature level, feature should have the same batch size and404            # spatial dimension as the box_cls and box_delta.405            dummy_features = [x.clone()[:, 0:0, :, :] for x in pred_logits]406            # self.num_classess can be inferred407            self.num_classes = pred_logits[0].shape[1] // (pred_anchor_deltas[0].shape[1] // 4)408 409            results = self.forward_inference(410                dummy_images, dummy_features, [pred_logits, pred_anchor_deltas]411            )412            return meta_arch.GeneralizedRCNN._postprocess(results, batched_inputs, image_sizes)413 414        return f415 416 417META_ARCH_CAFFE2_EXPORT_TYPE_MAP = {418    "GeneralizedRCNN": Caffe2GeneralizedRCNN,419    "RetinaNet": Caffe2RetinaNet,420}421