Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
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 