Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
vision_transformer_utils.py550 linesDownload Raw Back to models
1import math2import pathlib3import warnings4from types import FunctionType5from typing import Any, BinaryIO, List, Optional, Tuple, Union6 7import numpy as np8import torch9from PIL import Image, ImageColor, ImageDraw, ImageFont10 11__all__ = [12    "make_grid",13    "save_image",14    "draw_bounding_boxes",15    "draw_segmentation_masks",16    "draw_keypoints",17    "flow_to_image",18]19 20 21@torch.no_grad()22def make_grid(23    tensor: Union[torch.Tensor, List[torch.Tensor]],24    nrow: int = 8,25    padding: int = 2,26    normalize: bool = False,27    value_range: Optional[Tuple[int, int]] = None,28    scale_each: bool = False,29    pad_value: float = 0.0,30    **kwargs,31) -> torch.Tensor:32    """33    Make a grid of images.34 35    Args:36        tensor (Tensor or list): 4D mini-batch Tensor of shape (B x C x H x W)37            or a list of images all of the same size.38        nrow (int, optional): Number of images displayed in each row of the grid.39            The final grid size is ``(B / nrow, nrow)``. Default: ``8``.40        padding (int, optional): amount of padding. Default: ``2``.41        normalize (bool, optional): If True, shift the image to the range (0, 1),42            by the min and max values specified by ``value_range``. Default: ``False``.43        value_range (tuple, optional): tuple (min, max) where min and max are numbers,44            then these numbers are used to normalize the image. By default, min and max45            are computed from the tensor.46        range (tuple. optional):47            .. warning::48                This parameter was deprecated in ``0.12`` and will be removed in ``0.14``. Please use ``value_range``49                instead.50        scale_each (bool, optional): If ``True``, scale each image in the batch of51            images separately rather than the (min, max) over all images. Default: ``False``.52        pad_value (float, optional): Value for the padded pixels. Default: ``0``.53 54    Returns:55        grid (Tensor): the tensor containing grid of images.56    """57    if not torch.jit.is_scripting() and not torch.jit.is_tracing():58        _log_api_usage_once(make_grid)59    if not (torch.is_tensor(tensor) or (isinstance(tensor, list) and all(torch.is_tensor(t) for t in tensor))):60        raise TypeError(f"tensor or list of tensors expected, got {type(tensor)}")61 62    if "range" in kwargs.keys():63        warnings.warn(64            "The parameter 'range' is deprecated since 0.12 and will be removed in 0.14. "65            "Please use 'value_range' instead."66        )67        value_range = kwargs["range"]68 69    # if list of tensors, convert to a 4D mini-batch Tensor70    if isinstance(tensor, list):71        tensor = torch.stack(tensor, dim=0)72 73    if tensor.dim() == 2:  # single image H x W74        tensor = tensor.unsqueeze(0)75    if tensor.dim() == 3:  # single image76        if tensor.size(0) == 1:  # if single-channel, convert to 3-channel77            tensor = torch.cat((tensor, tensor, tensor), 0)78        tensor = tensor.unsqueeze(0)79 80    if tensor.dim() == 4 and tensor.size(1) == 1:  # single-channel images81        tensor = torch.cat((tensor, tensor, tensor), 1)82 83    if normalize is True:84        tensor = tensor.clone()  # avoid modifying tensor in-place85        if value_range is not None:86            assert isinstance(87                value_range, tuple88            ), "value_range has to be a tuple (min, max) if specified. min and max are numbers"89 90        def norm_ip(img, low, high):91            img.clamp_(min=low, max=high)92            img.sub_(low).div_(max(high - low, 1e-5))93 94        def norm_range(t, value_range):95            if value_range is not None:96                norm_ip(t, value_range[0], value_range[1])97            else:98                norm_ip(t, float(t.min()), float(t.max()))99 100        if scale_each is True:101            for t in tensor:  # loop over mini-batch dimension102                norm_range(t, value_range)103        else:104            norm_range(tensor, value_range)105 106    assert isinstance(tensor, torch.Tensor)107    if tensor.size(0) == 1:108        return tensor.squeeze(0)109 110    # make the mini-batch of images into a grid111    nmaps = tensor.size(0)112    xmaps = min(nrow, nmaps)113    ymaps = int(math.ceil(float(nmaps) / xmaps))114    height, width = int(tensor.size(2) + padding), int(tensor.size(3) + padding)115    num_channels = tensor.size(1)116    grid = tensor.new_full((num_channels, height * ymaps + padding, width * xmaps + padding), pad_value)117    k = 0118    for y in range(ymaps):119        for x in range(xmaps):120            if k >= nmaps:121                break122            # Tensor.copy_() is a valid method but seems to be missing from the stubs123            # https://pytorch.org/docs/stable/tensors.html#torch.Tensor.copy_124            grid.narrow(1, y * height + padding, height - padding).narrow(  # type: ignore[attr-defined]125                2, x * width + padding, width - padding126            ).copy_(tensor[k])127            k = k + 1128    return grid129 130 131@torch.no_grad()132def save_image(133    tensor: Union[torch.Tensor, List[torch.Tensor]],134    fp: Union[str, pathlib.Path, BinaryIO],135    format: Optional[str] = None,136    **kwargs,137) -> None:138    """139    Save a given Tensor into an image file.140 141    Args:142        tensor (Tensor or list): Image to be saved. If given a mini-batch tensor,143            saves the tensor as a grid of images by calling ``make_grid``.144        fp (string or file object): A filename or a file object145        format(Optional):  If omitted, the format to use is determined from the filename extension.146            If a file object was used instead of a filename, this parameter should always be used.147        **kwargs: Other arguments are documented in ``make_grid``.148    """149 150    if not torch.jit.is_scripting() and not torch.jit.is_tracing():151        _log_api_usage_once(save_image)152    grid = make_grid(tensor, **kwargs)153    # Add 0.5 after unnormalizing to [0, 255] to round to nearest integer154    ndarr = grid.mul(255).add_(0.5).clamp_(0, 255).permute(1, 2, 0).to("cpu", torch.uint8).numpy()155    im = Image.fromarray(ndarr)156    im.save(fp, format=format)157 158 159@torch.no_grad()160def draw_bounding_boxes(161    image: torch.Tensor,162    boxes: torch.Tensor,163    labels: Optional[List[str]] = None,164    colors: Optional[Union[List[Union[str, Tuple[int, int, int]]], str, Tuple[int, int, int]]] = None,165    fill: Optional[bool] = False,166    width: int = 1,167    font: Optional[str] = None,168    font_size: int = 10,169) -> torch.Tensor:170 171    """172    Draws bounding boxes on given image.173    The values of the input image should be uint8 between 0 and 255.174    If fill is True, Resulting Tensor should be saved as PNG image.175 176    Args:177        image (Tensor): Tensor of shape (C x H x W) and dtype uint8.178        boxes (Tensor): Tensor of size (N, 4) containing bounding boxes in (xmin, ymin, xmax, ymax) format. Note that179            the boxes are absolute coordinates with respect to the image. In other words: `0 <= xmin < xmax < W` and180            `0 <= ymin < ymax < H`.181        labels (List[str]): List containing the labels of bounding boxes.182        colors (color or list of colors, optional): List containing the colors183            of the boxes or single color for all boxes. The color can be represented as184            PIL strings e.g. "red" or "#FF00FF", or as RGB tuples e.g. ``(240, 10, 157)``.185            By default, random colors are generated for boxes.186        fill (bool): If `True` fills the bounding box with specified color.187        width (int): Width of bounding box.188        font (str): A filename containing a TrueType font. If the file is not found in this filename, the loader may189            also search in other directories, such as the `fonts/` directory on Windows or `/Library/Fonts/`,190            `/System/Library/Fonts/` and `~/Library/Fonts/` on macOS.191        font_size (int): The requested font size in points.192 193    Returns:194        img (Tensor[C, H, W]): Image Tensor of dtype uint8 with bounding boxes plotted.195    """196 197    if not torch.jit.is_scripting() and not torch.jit.is_tracing():198        _log_api_usage_once(draw_bounding_boxes)199    if not isinstance(image, torch.Tensor):200        raise TypeError(f"Tensor expected, got {type(image)}")201    elif image.dtype != torch.uint8:202        raise ValueError(f"Tensor uint8 expected, got {image.dtype}")203    elif image.dim() != 3:204        raise ValueError("Pass individual images, not batches")205    elif image.size(0) not in {1, 3}:206        raise ValueError("Only grayscale and RGB images are supported")207 208    num_boxes = boxes.shape[0]209 210    if labels is None:211        labels: Union[List[str], List[None]] = [None] * num_boxes  # type: ignore[no-redef]212    elif len(labels) != num_boxes:213        raise ValueError(214            f"Number of boxes ({num_boxes}) and labels ({len(labels)}) mismatch. Please specify labels for each box."215        )216 217    if colors is None:218        colors = _generate_color_palette(num_boxes)219    elif isinstance(colors, list):220        if len(colors) < num_boxes:221            raise ValueError(f"Number of colors ({len(colors)}) is less than number of boxes ({num_boxes}). ")222    else:  # colors specifies a single color for all boxes223        colors = [colors] * num_boxes224 225    colors = [(ImageColor.getrgb(color) if isinstance(color, str) else color) for color in colors]226 227    # Handle Grayscale images228    if image.size(0) == 1:229        image = torch.tile(image, (3, 1, 1))230 231    ndarr = image.permute(1, 2, 0).cpu().numpy()232    img_to_draw = Image.fromarray(ndarr)233    img_boxes = boxes.to(torch.int64).tolist()234 235    if fill:236        draw = ImageDraw.Draw(img_to_draw, "RGBA")237    else:238        draw = ImageDraw.Draw(img_to_draw)239 240    txt_font = ImageFont.load_default() if font is None else ImageFont.truetype(font=font, size=font_size)241 242    for bbox, color, label in zip(img_boxes, colors, labels):  # type: ignore[arg-type]243        if fill:244            fill_color = color + (100,)245            draw.rectangle(bbox, width=width, outline=color, fill=fill_color)246        else:247            draw.rectangle(bbox, width=width, outline=color)248 249        if label is not None:250            margin = width + 1251            draw.text((bbox[0] + margin, bbox[1] + margin), label, fill=color, font=txt_font)252 253    return torch.from_numpy(np.array(img_to_draw)).permute(2, 0, 1).to(dtype=torch.uint8)254 255 256@torch.no_grad()257def draw_segmentation_masks(258    image: torch.Tensor,259    masks: torch.Tensor,260    alpha: float = 0.8,261    colors: Optional[Union[List[Union[str, Tuple[int, int, int]]], str, Tuple[int, int, int]]] = None,262) -> torch.Tensor:263 264    """265    Draws segmentation masks on given RGB image.266    The values of the input image should be uint8 between 0 and 255.267 268    Args:269        image (Tensor): Tensor of shape (3, H, W) and dtype uint8.270        masks (Tensor): Tensor of shape (num_masks, H, W) or (H, W) and dtype bool.271        alpha (float): Float number between 0 and 1 denoting the transparency of the masks.272            0 means full transparency, 1 means no transparency.273        colors (color or list of colors, optional): List containing the colors274            of the masks or single color for all masks. The color can be represented as275            PIL strings e.g. "red" or "#FF00FF", or as RGB tuples e.g. ``(240, 10, 157)``.276            By default, random colors are generated for each mask.277 278    Returns:279        img (Tensor[C, H, W]): Image Tensor, with segmentation masks drawn on top.280    """281 282    if not torch.jit.is_scripting() and not torch.jit.is_tracing():283        _log_api_usage_once(draw_segmentation_masks)284    if not isinstance(image, torch.Tensor):285        raise TypeError(f"The image must be a tensor, got {type(image)}")286    elif image.dtype != torch.uint8:287        raise ValueError(f"The image dtype must be uint8, got {image.dtype}")288    elif image.dim() != 3:289        raise ValueError("Pass individual images, not batches")290    elif image.size()[0] != 3:291        raise ValueError("Pass an RGB image. Other Image formats are not supported")292    if masks.ndim == 2:293        masks = masks[None, :, :]294    if masks.ndim != 3:295        raise ValueError("masks must be of shape (H, W) or (batch_size, H, W)")296    if masks.dtype != torch.bool:297        raise ValueError(f"The masks must be of dtype bool. Got {masks.dtype}")298    if masks.shape[-2:] != image.shape[-2:]:299        raise ValueError("The image and the masks must have the same height and width")300 301    num_masks = masks.size()[0]302    if colors is not None and num_masks > len(colors):303        raise ValueError(f"There are more masks ({num_masks}) than colors ({len(colors)})")304 305    if colors is None:306        colors = _generate_color_palette(num_masks)307 308    if not isinstance(colors, list):309        colors = [colors]310    if not isinstance(colors[0], (tuple, str)):311        raise ValueError("colors must be a tuple or a string, or a list thereof")312    if isinstance(colors[0], tuple) and len(colors[0]) != 3:313        raise ValueError("It seems that you passed a tuple of colors instead of a list of colors")314 315    out_dtype = torch.uint8316 317    colors_ = []318    for color in colors:319        if isinstance(color, str):320            color = ImageColor.getrgb(color)321        colors_.append(torch.tensor(color, dtype=out_dtype))322 323    img_to_draw = image.detach().clone()324    # TODO: There might be a way to vectorize this325    for mask, color in zip(masks, colors_):326        img_to_draw[:, mask] = color[:, None]327 328    out = image * (1 - alpha) + img_to_draw * alpha329    return out.to(out_dtype)330 331 332@torch.no_grad()333def draw_keypoints(334    image: torch.Tensor,335    keypoints: torch.Tensor,336    connectivity: Optional[List[Tuple[int, int]]] = None,337    colors: Optional[Union[str, Tuple[int, int, int]]] = None,338    radius: int = 2,339    width: int = 3,340) -> torch.Tensor:341 342    """343    Draws Keypoints on given RGB image.344    The values of the input image should be uint8 between 0 and 255.345 346    Args:347        image (Tensor): Tensor of shape (3, H, W) and dtype uint8.348        keypoints (Tensor): Tensor of shape (num_instances, K, 2) the K keypoints location for each of the N instances,349            in the format [x, y].350        connectivity (List[Tuple[int, int]]]): A List of tuple where,351            each tuple contains pair of keypoints to be connected.352        colors (str, Tuple): The color can be represented as353            PIL strings e.g. "red" or "#FF00FF", or as RGB tuples e.g. ``(240, 10, 157)``.354        radius (int): Integer denoting radius of keypoint.355        width (int): Integer denoting width of line connecting keypoints.356 357    Returns:358        img (Tensor[C, H, W]): Image Tensor of dtype uint8 with keypoints drawn.359    """360 361    if not torch.jit.is_scripting() and not torch.jit.is_tracing():362        _log_api_usage_once(draw_keypoints)363    if not isinstance(image, torch.Tensor):364        raise TypeError(f"The image must be a tensor, got {type(image)}")365    elif image.dtype != torch.uint8:366        raise ValueError(f"The image dtype must be uint8, got {image.dtype}")367    elif image.dim() != 3:368        raise ValueError("Pass individual images, not batches")369    elif image.size()[0] != 3:370        raise ValueError("Pass an RGB image. Other Image formats are not supported")371 372    if keypoints.ndim != 3:373        raise ValueError("keypoints must be of shape (num_instances, K, 2)")374 375    ndarr = image.permute(1, 2, 0).cpu().numpy()376    img_to_draw = Image.fromarray(ndarr)377    draw = ImageDraw.Draw(img_to_draw)378    img_kpts = keypoints.to(torch.int64).tolist()379 380    for kpt_id, kpt_inst in enumerate(img_kpts):381        for inst_id, kpt in enumerate(kpt_inst):382            x1 = kpt[0] - radius383            x2 = kpt[0] + radius384            y1 = kpt[1] - radius385            y2 = kpt[1] + radius386            draw.ellipse([x1, y1, x2, y2], fill=colors, outline=None, width=0)387 388        if connectivity:389            for connection in connectivity:390                start_pt_x = kpt_inst[connection[0]][0]391                start_pt_y = kpt_inst[connection[0]][1]392 393                end_pt_x = kpt_inst[connection[1]][0]394                end_pt_y = kpt_inst[connection[1]][1]395 396                draw.line(397                    ((start_pt_x, start_pt_y), (end_pt_x, end_pt_y)),398                    width=width,399                )400 401    return torch.from_numpy(np.array(img_to_draw)).permute(2, 0, 1).to(dtype=torch.uint8)402 403 404# Flow visualization code adapted from https://github.com/tomrunia/OpticalFlow_Visualization405@torch.no_grad()406def flow_to_image(flow: torch.Tensor) -> torch.Tensor:407 408    """409    Converts a flow to an RGB image.410 411    Args:412        flow (Tensor): Flow of shape (N, 2, H, W) or (2, H, W) and dtype torch.float.413 414    Returns:415        img (Tensor): Image Tensor of dtype uint8 where each color corresponds416            to a given flow direction. Shape is (N, 3, H, W) or (3, H, W) depending on the input.417    """418 419    if flow.dtype != torch.float:420        raise ValueError(f"Flow should be of dtype torch.float, got {flow.dtype}.")421 422    orig_shape = flow.shape423    if flow.ndim == 3:424        flow = flow[None]  # Add batch dim425 426    if flow.ndim != 4 or flow.shape[1] != 2:427        raise ValueError(f"Input flow should have shape (2, H, W) or (N, 2, H, W), got {orig_shape}.")428 429    max_norm = torch.sum(flow ** 2, dim=1).sqrt().max()430    epsilon = torch.finfo((flow).dtype).eps431    normalized_flow = flow / (max_norm + epsilon)432    img = _normalized_flow_to_image(normalized_flow)433 434    if len(orig_shape) == 3:435        img = img[0]  # Remove batch dim436    return img437 438 439@torch.no_grad()440def _normalized_flow_to_image(normalized_flow: torch.Tensor) -> torch.Tensor:441 442    """443    Converts a batch of normalized flow to an RGB image.444 445    Args:446        normalized_flow (torch.Tensor): Normalized flow tensor of shape (N, 2, H, W)447    Returns:448       img (Tensor(N, 3, H, W)): Flow visualization image of dtype uint8.449    """450 451    N, _, H, W = normalized_flow.shape452    device = normalized_flow.device453    flow_image = torch.zeros((N, 3, H, W), dtype=torch.uint8, device=device)454    colorwheel = _make_colorwheel().to(device)  # shape [55x3]455    num_cols = colorwheel.shape[0]456    norm = torch.sum(normalized_flow ** 2, dim=1).sqrt()457    a = torch.atan2(-normalized_flow[:, 1, :, :], -normalized_flow[:, 0, :, :]) / torch.pi458    fk = (a + 1) / 2 * (num_cols - 1)459    k0 = torch.floor(fk).to(torch.long)460    k1 = k0 + 1461    k1[k1 == num_cols] = 0462    f = fk - k0463 464    for c in range(colorwheel.shape[1]):465        tmp = colorwheel[:, c]466        col0 = tmp[k0] / 255.0467        col1 = tmp[k1] / 255.0468        col = (1 - f) * col0 + f * col1469        col = 1 - norm * (1 - col)470        flow_image[:, c, :, :] = torch.floor(255 * col)471    return flow_image472 473 474def _make_colorwheel() -> torch.Tensor:475    """476    Generates a color wheel for optical flow visualization as presented in:477    Baker et al. "A Database and Evaluation Methodology for Optical Flow" (ICCV, 2007)478    URL: http://vision.middlebury.edu/flow/flowEval-iccv07.pdf.479 480    Returns:481        colorwheel (Tensor[55, 3]): Colorwheel Tensor.482    """483 484    RY = 15485    YG = 6486    GC = 4487    CB = 11488    BM = 13489    MR = 6490 491    ncols = RY + YG + GC + CB + BM + MR492    colorwheel = torch.zeros((ncols, 3))493    col = 0494 495    # RY496    colorwheel[0:RY, 0] = 255497    colorwheel[0:RY, 1] = torch.floor(255 * torch.arange(0, RY) / RY)498    col = col + RY499    # YG500    colorwheel[col : col + YG, 0] = 255 - torch.floor(255 * torch.arange(0, YG) / YG)501    colorwheel[col : col + YG, 1] = 255502    col = col + YG503    # GC504    colorwheel[col : col + GC, 1] = 255505    colorwheel[col : col + GC, 2] = torch.floor(255 * torch.arange(0, GC) / GC)506    col = col + GC507    # CB508    colorwheel[col : col + CB, 1] = 255 - torch.floor(255 * torch.arange(CB) / CB)509    colorwheel[col : col + CB, 2] = 255510    col = col + CB511    # BM512    colorwheel[col : col + BM, 2] = 255513    colorwheel[col : col + BM, 0] = torch.floor(255 * torch.arange(0, BM) / BM)514    col = col + BM515    # MR516    colorwheel[col : col + MR, 2] = 255 - torch.floor(255 * torch.arange(MR) / MR)517    colorwheel[col : col + MR, 0] = 255518    return colorwheel519 520 521def _generate_color_palette(num_objects: int):522    palette = torch.tensor([2 ** 25 - 1, 2 ** 15 - 1, 2 ** 21 - 1])523    return [tuple((i * palette) % 255) for i in range(num_objects)]524 525 526def _log_api_usage_once(obj: Any) -> None:527 528    """529    Logs API usage(module and name) within an organization.530    In a large ecosystem, it's often useful to track the PyTorch and531    TorchVision APIs usage. This API provides the similar functionality to the532    logging module in the Python stdlib. It can be used for debugging purpose533    to log which methods are used and by default it is inactive, unless the user534    manually subscribes a logger via the `SetAPIUsageLogger method <https://github.com/pytorch/pytorch/blob/eb3b9fe719b21fae13c7a7cf3253f970290a573e/c10/util/Logging.cpp#L114>`_.535    Please note it is triggered only once for the same API call within a process.536    It does not collect any data from open-source users since it is no-op by default.537    For more information, please refer to538    * PyTorch note: https://pytorch.org/docs/stable/notes/large_scale_deployments.html#api-usage-logging;539    * Logging policy: https://github.com/pytorch/vision/issues/5052;540 541    Args:542        obj (class instance or method): an object to extract info from.543    """544    if not obj.__module__.startswith("torchvision"):545        return546    name = obj.__class__.__name__547    if isinstance(obj, FunctionType):548        name = obj.__name__549    torch._C._log_api_usage_once(f"{obj.__module__}.{name}")550