David310/Detect_AI-generated_Image
4
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 