Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

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