Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
keypoints.py236 linesDownload Raw Back to structures
1# Copyright (c) Facebook, Inc. and its affiliates.2import numpy as np3from typing import Any, List, Tuple, Union4import torch5from torch.nn import functional as F6 7 8class Keypoints:9    """10    Stores keypoint **annotation** data. GT Instances have a `gt_keypoints` property11    containing the x,y location and visibility flag of each keypoint. This tensor has shape12    (N, K, 3) where N is the number of instances and K is the number of keypoints per instance.13 14    The visibility flag follows the COCO format and must be one of three integers:15 16    * v=0: not labeled (in which case x=y=0)17    * v=1: labeled but not visible18    * v=2: labeled and visible19    """20 21    def __init__(self, keypoints: Union[torch.Tensor, np.ndarray, List[List[float]]]):22        """23        Arguments:24            keypoints: A Tensor, numpy array, or list of the x, y, and visibility of each keypoint.25                The shape should be (N, K, 3) where N is the number of26                instances, and K is the number of keypoints per instance.27        """28        device = keypoints.device if isinstance(keypoints, torch.Tensor) else torch.device("cpu")29        keypoints = torch.as_tensor(keypoints, dtype=torch.float32, device=device)30        assert keypoints.dim() == 3 and keypoints.shape[2] == 3, keypoints.shape31        self.tensor = keypoints32 33    def __len__(self) -> int:34        return self.tensor.size(0)35 36    def to(self, *args: Any, **kwargs: Any) -> "Keypoints":37        return type(self)(self.tensor.to(*args, **kwargs))38 39    @property40    def device(self) -> torch.device:41        return self.tensor.device42 43    def to_heatmap(self, boxes: torch.Tensor, heatmap_size: int) -> torch.Tensor:44        """45        Convert keypoint annotations to a heatmap of one-hot labels for training,46        as described in :paper:`Mask R-CNN`.47 48        Arguments:49            boxes: Nx4 tensor, the boxes to draw the keypoints to50 51        Returns:52            heatmaps:53                A tensor of shape (N, K), each element is integer spatial label54                in the range [0, heatmap_size**2 - 1] for each keypoint in the input.55            valid:56                A tensor of shape (N, K) containing whether each keypoint is in the roi or not.57        """58        return _keypoints_to_heatmap(self.tensor, boxes, heatmap_size)59 60    def __getitem__(self, item: Union[int, slice, torch.BoolTensor]) -> "Keypoints":61        """62        Create a new `Keypoints` by indexing on this `Keypoints`.63 64        The following usage are allowed:65 66        1. `new_kpts = kpts[3]`: return a `Keypoints` which contains only one instance.67        2. `new_kpts = kpts[2:10]`: return a slice of key points.68        3. `new_kpts = kpts[vector]`, where vector is a torch.ByteTensor69           with `length = len(kpts)`. Nonzero elements in the vector will be selected.70 71        Note that the returned Keypoints might share storage with this Keypoints,72        subject to Pytorch's indexing semantics.73        """74        if isinstance(item, int):75            return Keypoints([self.tensor[item]])76        return Keypoints(self.tensor[item])77 78    def __repr__(self) -> str:79        s = self.__class__.__name__ + "("80        s += "num_instances={})".format(len(self.tensor))81        return s82 83    @staticmethod84    def cat(keypoints_list: List["Keypoints"]) -> "Keypoints":85        """86        Concatenates a list of Keypoints into a single Keypoints87 88        Arguments:89            keypoints_list (list[Keypoints])90 91        Returns:92            Keypoints: the concatenated Keypoints93        """94        assert isinstance(keypoints_list, (list, tuple))95        assert len(keypoints_list) > 096        assert all(isinstance(keypoints, Keypoints) for keypoints in keypoints_list)97 98        cat_kpts = type(keypoints_list[0])(99            torch.cat([kpts.tensor for kpts in keypoints_list], dim=0)100        )101        return cat_kpts102 103 104# TODO make this nicer, this is a direct translation from C2 (but removing the inner loop)105def _keypoints_to_heatmap(106    keypoints: torch.Tensor, rois: torch.Tensor, heatmap_size: int107) -> Tuple[torch.Tensor, torch.Tensor]:108    """109    Encode keypoint locations into a target heatmap for use in SoftmaxWithLoss across space.110 111    Maps keypoints from the half-open interval [x1, x2) on continuous image coordinates to the112    closed interval [0, heatmap_size - 1] on discrete image coordinates. We use the113    continuous-discrete conversion from Heckbert 1990 ("What is the coordinate of a pixel?"):114    d = floor(c) and c = d + 0.5, where d is a discrete coordinate and c is a continuous coordinate.115 116    Arguments:117        keypoints: tensor of keypoint locations in of shape (N, K, 3).118        rois: Nx4 tensor of rois in xyxy format119        heatmap_size: integer side length of square heatmap.120 121    Returns:122        heatmaps: A tensor of shape (N, K) containing an integer spatial label123            in the range [0, heatmap_size**2 - 1] for each keypoint in the input.124        valid: A tensor of shape (N, K) containing whether each keypoint is in125            the roi or not.126    """127 128    if rois.numel() == 0:129        return rois.new().long(), rois.new().long()130    offset_x = rois[:, 0]131    offset_y = rois[:, 1]132    scale_x = heatmap_size / (rois[:, 2] - rois[:, 0])133    scale_y = heatmap_size / (rois[:, 3] - rois[:, 1])134 135    offset_x = offset_x[:, None]136    offset_y = offset_y[:, None]137    scale_x = scale_x[:, None]138    scale_y = scale_y[:, None]139 140    x = keypoints[..., 0]141    y = keypoints[..., 1]142 143    x_boundary_inds = x == rois[:, 2][:, None]144    y_boundary_inds = y == rois[:, 3][:, None]145 146    x = (x - offset_x) * scale_x147    x = x.floor().long()148    y = (y - offset_y) * scale_y149    y = y.floor().long()150 151    x[x_boundary_inds] = heatmap_size - 1152    y[y_boundary_inds] = heatmap_size - 1153 154    valid_loc = (x >= 0) & (y >= 0) & (x < heatmap_size) & (y < heatmap_size)155    vis = keypoints[..., 2] > 0156    valid = (valid_loc & vis).long()157 158    lin_ind = y * heatmap_size + x159    heatmaps = lin_ind * valid160 161    return heatmaps, valid162 163 164@torch.jit.script_if_tracing165def heatmaps_to_keypoints(maps: torch.Tensor, rois: torch.Tensor) -> torch.Tensor:166    """167    Extract predicted keypoint locations from heatmaps.168 169    Args:170        maps (Tensor): (#ROIs, #keypoints, POOL_H, POOL_W). The predicted heatmap of logits for171            each ROI and each keypoint.172        rois (Tensor): (#ROIs, 4). The box of each ROI.173 174    Returns:175        Tensor of shape (#ROIs, #keypoints, 4) with the last dimension corresponding to176        (x, y, logit, score) for each keypoint.177 178    When converting discrete pixel indices in an NxN image to a continuous keypoint coordinate,179    we maintain consistency with :meth:`Keypoints.to_heatmap` by using the conversion from180    Heckbert 1990: c = d + 0.5, where d is a discrete coordinate and c is a continuous coordinate.181    """182 183    offset_x = rois[:, 0]184    offset_y = rois[:, 1]185 186    widths = (rois[:, 2] - rois[:, 0]).clamp(min=1)187    heights = (rois[:, 3] - rois[:, 1]).clamp(min=1)188    widths_ceil = widths.ceil()189    heights_ceil = heights.ceil()190 191    num_rois, num_keypoints = maps.shape[:2]192    xy_preds = maps.new_zeros(rois.shape[0], num_keypoints, 4)193 194    width_corrections = widths / widths_ceil195    height_corrections = heights / heights_ceil196 197    keypoints_idx = torch.arange(num_keypoints, device=maps.device)198 199    for i in range(num_rois):200        outsize = (int(heights_ceil[i]), int(widths_ceil[i]))201        roi_map = F.interpolate(maps[[i]], size=outsize, mode="bicubic", align_corners=False)202 203        # Although semantically equivalent, `reshape` is used instead of `squeeze` due204        # to limitation during ONNX export of `squeeze` in scripting mode205        roi_map = roi_map.reshape(roi_map.shape[1:])  # keypoints x H x W206 207        # softmax over the spatial region208        max_score, _ = roi_map.view(num_keypoints, -1).max(1)209        max_score = max_score.view(num_keypoints, 1, 1)210        tmp_full_resolution = (roi_map - max_score).exp_()211        tmp_pool_resolution = (maps[i] - max_score).exp_()212        # Produce scores over the region H x W, but normalize with POOL_H x POOL_W,213        # so that the scores of objects of different absolute sizes will be more comparable214        roi_map_scores = tmp_full_resolution / tmp_pool_resolution.sum((1, 2), keepdim=True)215 216        w = roi_map.shape[2]217        pos = roi_map.view(num_keypoints, -1).argmax(1)218 219        x_int = pos % w220        y_int = (pos - x_int) // w221 222        assert (223            roi_map_scores[keypoints_idx, y_int, x_int]224            == roi_map_scores.view(num_keypoints, -1).max(1)[0]225        ).all()226 227        x = (x_int.float() + 0.5) * width_corrections[i]228        y = (y_int.float() + 0.5) * height_corrections[i]229 230        xy_preds[i, :, 0] = x + offset_x[i]231        xy_preds[i, :, 1] = y + offset_y[i]232        xy_preds[i, :, 2] = roi_map[keypoints_idx, y_int, x_int]233        xy_preds[i, :, 3] = roi_map_scores[keypoints_idx, y_int, x_int]234 235    return xy_preds236