Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
mask_ops.py276 linesDownload Raw Back to layers
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