Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import numpy as np3from typing import Tuple4import torch5from PIL import Image6from torch.nn import functional as F7 8__all__ = ["paste_masks_in_image"]9 10 11BYTES_PER_FLOAT = 412# TODO: This memory limit may be too much or too little. It would be better to13# determine it based on available resources.14GPU_MEM_LIMIT = 1024**3 # 1 GB memory limit15 16 17def _do_paste_mask(masks, boxes, img_h: int, img_w: int, skip_empty: bool = True):18 """19 Args:20 masks: N, 1, H, W21 boxes: N, 422 img_h, img_w (int):23 skip_empty (bool): only paste masks within the region that24 tightly bound all boxes, and returns the results this region only.25 An important optimization for CPU.26 27 Returns:28 if skip_empty == False, a mask of shape (N, img_h, img_w)29 if skip_empty == True, a mask of shape (N, h', w'), and the slice30 object for the corresponding region.31 """32 # On GPU, paste all masks together (up to chunk size)33 # by using the entire image to sample the masks34 # Compared to pasting them one by one,35 # this has more operations but is faster on COCO-scale dataset.36 device = masks.device37 38 if skip_empty and not torch.jit.is_scripting():39 x0_int, y0_int = torch.clamp(boxes.min(dim=0).values.floor()[:2] - 1, min=0).to(40 dtype=torch.int3241 )42 x1_int = torch.clamp(boxes[:, 2].max().ceil() + 1, max=img_w).to(dtype=torch.int32)43 y1_int = torch.clamp(boxes[:, 3].max().ceil() + 1, max=img_h).to(dtype=torch.int32)44 else:45 x0_int, y0_int = 0, 046 x1_int, y1_int = img_w, img_h47 x0, y0, x1, y1 = torch.split(boxes, 1, dim=1) # each is Nx148 49 N = masks.shape[0]50 51 img_y = torch.arange(y0_int, y1_int, device=device, dtype=torch.float32) + 0.552 img_x = torch.arange(x0_int, x1_int, device=device, dtype=torch.float32) + 0.553 img_y = (img_y - y0) / (y1 - y0) * 2 - 154 img_x = (img_x - x0) / (x1 - x0) * 2 - 155 # img_x, img_y have shapes (N, w), (N, h)56 57 gx = img_x[:, None, :].expand(N, img_y.size(1), img_x.size(1))58 gy = img_y[:, :, None].expand(N, img_y.size(1), img_x.size(1))59 grid = torch.stack([gx, gy], dim=3)60 61 if not torch.jit.is_scripting():62 if not masks.dtype.is_floating_point:63 masks = masks.float()64 img_masks = F.grid_sample(masks, grid.to(masks.dtype), align_corners=False)65 66 if skip_empty and not torch.jit.is_scripting():67 return img_masks[:, 0], (slice(y0_int, y1_int), slice(x0_int, x1_int))68 else:69 return img_masks[:, 0], ()70 71 72# Annotate boxes as Tensor (but not Boxes) in order to use scripting73@torch.jit.script_if_tracing74def paste_masks_in_image(75 masks: torch.Tensor, boxes: torch.Tensor, image_shape: Tuple[int, int], threshold: float = 0.576):77 """78 Paste a set of masks that are of a fixed resolution (e.g., 28 x 28) into an image.79 The location, height, and width for pasting each mask is determined by their80 corresponding bounding boxes in boxes.81 82 Note:83 This is a complicated but more accurate implementation. In actual deployment, it is84 often enough to use a faster but less accurate implementation.85 See :func:`paste_mask_in_image_old` in this file for an alternative implementation.86 87 Args:88 masks (tensor): Tensor of shape (Bimg, Hmask, Wmask), where Bimg is the number of89 detected object instances in the image and Hmask, Wmask are the mask width and mask90 height of the predicted mask (e.g., Hmask = Wmask = 28). Values are in [0, 1].91 boxes (Boxes or Tensor): A Boxes of length Bimg or Tensor of shape (Bimg, 4).92 boxes[i] and masks[i] correspond to the same object instance.93 image_shape (tuple): height, width94 threshold (float): A threshold in [0, 1] for converting the (soft) masks to95 binary masks.96 97 Returns:98 img_masks (Tensor): A tensor of shape (Bimg, Himage, Wimage), where Bimg is the99 number of detected object instances and Himage, Wimage are the image width100 and height. img_masks[i] is a binary mask for object instance i.101 """102 103 assert masks.shape[-1] == masks.shape[-2], "Only square mask predictions are supported"104 N = len(masks)105 if N == 0:106 return masks.new_empty((0,) + image_shape, dtype=torch.uint8)107 if not isinstance(boxes, torch.Tensor):108 boxes = boxes.tensor109 device = boxes.device110 assert len(boxes) == N, boxes.shape111 112 img_h, img_w = image_shape113 114 # The actual implementation split the input into chunks,115 # and paste them chunk by chunk.116 if device.type == "cpu" or torch.jit.is_scripting():117 # CPU is most efficient when they are pasted one by one with skip_empty=True118 # so that it performs minimal number of operations.119 num_chunks = N120 else:121 # GPU benefits from parallelism for larger chunks, but may have memory issue122 # int(img_h) because shape may be tensors in tracing123 num_chunks = int(np.ceil(N * int(img_h) * int(img_w) * BYTES_PER_FLOAT / GPU_MEM_LIMIT))124 assert (125 num_chunks <= N126 ), "Default GPU_MEM_LIMIT in mask_ops.py is too small; try increasing it"127 chunks = torch.chunk(torch.arange(N, device=device), num_chunks)128 129 img_masks = torch.zeros(130 N, img_h, img_w, device=device, dtype=torch.bool if threshold >= 0 else torch.uint8131 )132 for inds in chunks:133 masks_chunk, spatial_inds = _do_paste_mask(134 masks[inds, None, :, :], boxes[inds], img_h, img_w, skip_empty=device.type == "cpu"135 )136 137 if threshold >= 0:138 masks_chunk = (masks_chunk >= threshold).to(dtype=torch.bool)139 else:140 # for visualization and debugging141 masks_chunk = (masks_chunk * 255).to(dtype=torch.uint8)142 143 if torch.jit.is_scripting(): # Scripting does not use the optimized codepath144 img_masks[inds] = masks_chunk145 else:146 img_masks[(inds,) + spatial_inds] = masks_chunk147 return img_masks148 149 150# The below are the original paste function (from Detectron1) which has151# larger quantization error.152# It is faster on CPU, while the aligned one is faster on GPU thanks to grid_sample.153 154 155def paste_mask_in_image_old(mask, box, img_h, img_w, threshold):156 """157 Paste a single mask in an image.158 This is a per-box implementation of :func:`paste_masks_in_image`.159 This function has larger quantization error due to incorrect pixel160 modeling and is not used any more.161 162 Args:163 mask (Tensor): A tensor of shape (Hmask, Wmask) storing the mask of a single164 object instance. Values are in [0, 1].165 box (Tensor): A tensor of shape (4, ) storing the x0, y0, x1, y1 box corners166 of the object instance.167 img_h, img_w (int): Image height and width.168 threshold (float): Mask binarization threshold in [0, 1].169 170 Returns:171 im_mask (Tensor):172 The resized and binarized object mask pasted into the original173 image plane (a tensor of shape (img_h, img_w)).174 """175 # Conversion from continuous box coordinates to discrete pixel coordinates176 # via truncation (cast to int32). This determines which pixels to paste the177 # mask onto.178 box = box.to(dtype=torch.int32) # Continuous to discrete coordinate conversion179 # An example (1D) box with continuous coordinates (x0=0.7, x1=4.3) will map to180 # a discrete coordinates (x0=0, x1=4). Note that box is mapped to 5 = x1 - x0 + 1181 # pixels (not x1 - x0 pixels).182 samples_w = box[2] - box[0] + 1 # Number of pixel samples, *not* geometric width183 samples_h = box[3] - box[1] + 1 # Number of pixel samples, *not* geometric height184 185 # Resample the mask from it's original grid to the new samples_w x samples_h grid186 mask = Image.fromarray(mask.cpu().numpy())187 mask = mask.resize((samples_w, samples_h), resample=Image.BILINEAR)188 mask = np.array(mask, copy=False)189 190 if threshold >= 0:191 mask = np.array(mask > threshold, dtype=np.uint8)192 mask = torch.from_numpy(mask)193 else:194 # for visualization and debugging, we also195 # allow it to return an unmodified mask196 mask = torch.from_numpy(mask * 255).to(torch.uint8)197 198 im_mask = torch.zeros((img_h, img_w), dtype=torch.uint8)199 x_0 = max(box[0], 0)200 x_1 = min(box[2] + 1, img_w)201 y_0 = max(box[1], 0)202 y_1 = min(box[3] + 1, img_h)203 204 im_mask[y_0:y_1, x_0:x_1] = mask[205 (y_0 - box[1]) : (y_1 - box[1]), (x_0 - box[0]) : (x_1 - box[0])206 ]207 return im_mask208 209 210# Our pixel modeling requires extrapolation for any continuous211# coordinate < 0.5 or > length - 0.5. When sampling pixels on the masks,212# we would like this extrapolation to be an interpolation between boundary values and zero,213# instead of using absolute zero or boundary values.214# Therefore `paste_mask_in_image_old` is often used with zero padding around the masks like this:215# masks, scale = pad_masks(masks[:, 0, :, :], 1)216# boxes = scale_boxes(boxes.tensor, scale)217 218 219def pad_masks(masks, padding):220 """221 Args:222 masks (tensor): A tensor of shape (B, M, M) representing B masks.223 padding (int): Number of cells to pad on all sides.224 225 Returns:226 The padded masks and the scale factor of the padding size / original size.227 """228 B = masks.shape[0]229 M = masks.shape[-1]230 pad2 = 2 * padding231 scale = float(M + pad2) / M232 padded_masks = masks.new_zeros((B, M + pad2, M + pad2))233 padded_masks[:, padding:-padding, padding:-padding] = masks234 return padded_masks, scale235 236 237def scale_boxes(boxes, scale):238 """239 Args:240 boxes (tensor): A tensor of shape (B, 4) representing B boxes with 4241 coords representing the corners x0, y0, x1, y1,242 scale (float): The box scaling factor.243 244 Returns:245 Scaled boxes.246 """247 w_half = (boxes[:, 2] - boxes[:, 0]) * 0.5248 h_half = (boxes[:, 3] - boxes[:, 1]) * 0.5249 x_c = (boxes[:, 2] + boxes[:, 0]) * 0.5250 y_c = (boxes[:, 3] + boxes[:, 1]) * 0.5251 252 w_half *= scale253 h_half *= scale254 255 scaled_boxes = torch.zeros_like(boxes)256 scaled_boxes[:, 0] = x_c - w_half257 scaled_boxes[:, 2] = x_c + w_half258 scaled_boxes[:, 1] = y_c - h_half259 scaled_boxes[:, 3] = y_c + h_half260 return scaled_boxes261 262 263@torch.jit.script_if_tracing264def _paste_masks_tensor_shape(265 masks: torch.Tensor,266 boxes: torch.Tensor,267 image_shape: Tuple[torch.Tensor, torch.Tensor],268 threshold: float = 0.5,269):270 """271 A wrapper of paste_masks_in_image where image_shape is Tensor.272 During tracing, shapes might be tensors instead of ints. The Tensor->int273 conversion should be scripted rather than traced.274 """275 return paste_masks_in_image(masks, boxes, (int(image_shape[0]), int(image_shape[1])), threshold)276 