# SPDX-License-Identifier: Apache-2.0 """Image inputs -> the 8-bit plane the model reads (host only, no ttnn). SuperPoint reads one channel. Like the Hugging Face ``SuperPointForKeypointDetection`` (and the HTTP server of this repo), the model reads channel 0 (R) of the RGB image: for a grayscale image that is the gray plane. The plane is then resized to the 480x640 network frame (bilinear, bit-exact with Pillow; on the device for the precompiled source sizes). """ from __future__ import annotations import io import os from typing import Any, List import numpy as np #: PIL modes whose band 0 is already R of ``convert("RGB")`` (no conversion needed). _R_FIRST_MODES = ("RGB", "RGBA", "RGBX", "L") def is_single_image(x: Any) -> bool: """True for one image; False for a list / tuple of images or a 4-D batch array.""" if isinstance(x, (list, tuple)): return False nd = getattr(x, "ndim", None) if nd == 4: return False return True def split_batch(x: Any) -> List[Any]: """A list / tuple of images, or a 4-D array / tensor (B, ...) -> a list of single images.""" if isinstance(x, (list, tuple)): return list(x) return [x[i] for i in range(x.shape[0])] def _pil_plane(im) -> np.ndarray: im.load() if im.mode in _R_FIRST_MODES: band = im if im.mode == "L" else im.getchannel(0) else: band = im.convert("RGB").getchannel(0) return np.asarray(band) def _array_plane(a: np.ndarray, bgr: bool, channels_first: bool) -> np.ndarray: if a.ndim == 3: if channels_first and a.shape[0] in (1, 3, 4): a = a[2 if (bgr and a.shape[0] >= 3) else 0] elif a.shape[-1] in (1, 3, 4): a = a[..., 2 if (bgr and a.shape[-1] >= 3) else 0] elif a.shape[0] in (1, 3, 4): a = a[2 if (bgr and a.shape[0] >= 3) else 0] else: raise ValueError(f"cannot find the channel axis of an image array of shape {a.shape}") elif a.ndim != 2: raise ValueError(f"expected an image array with 2 or 3 dimensions, got shape {a.shape}") if a.dtype == np.uint8: return a if a.dtype == np.bool_: raise TypeError("boolean image arrays are not supported") if np.issubdtype(a.dtype, np.floating): # float images in [0, 1] (e.g. HF pixel_values / torchvision ToTensor): the device takes # 8-bit images, so the values are rounded to the nearest k/255. if a.size and (np.nanmin(a) < 0.0 or np.nanmax(a) > 1.0): raise ValueError("float images must have values in [0, 1]") return np.rint(np.nan_to_num(a.astype(np.float32)) * 255.0).astype(np.uint8) if np.issubdtype(a.dtype, np.integer): if a.size and (a.min() < 0 or a.max() > 255): raise ValueError("integer images must have values in 0..255") return a.astype(np.uint8) raise TypeError(f"unsupported image dtype {a.dtype}") def to_plane(image: Any, *, bgr: bool = False) -> np.ndarray: """One image -> the (H, W) uint8 plane the model reads, at the image's own size. Accepted inputs: * ``str`` / ``os.PathLike``: an image file (PNG, JPEG, ... anything Pillow opens). * ``bytes``: the encoded file contents. * ``PIL.Image.Image``: any mode; the R band of ``convert("RGB")`` is used. * ``numpy.ndarray``: (H, W) gray, or (H, W, C) with C in 1/3/4 (RGB / RGBA; BGR with ``bgr=True``, e.g. ``cv2.imread``). (C, H, W) is accepted when the last axis is not a channel axis. * ``torch.Tensor``: (H, W), (C, H, W) (torchvision convention) or (H, W, C). dtypes: uint8 (0..255), other integers in 0..255, or floats in [0, 1] (rounded to 8 bits). """ if isinstance(image, (str, os.PathLike)): from PIL import Image with Image.open(image) as im: return _pil_plane(im) if isinstance(image, (bytes, bytearray, memoryview)): from PIL import Image with Image.open(io.BytesIO(bytes(image))) as im: return _pil_plane(im) if hasattr(image, "getchannel") and hasattr(image, "mode"): # PIL.Image.Image return _pil_plane(image) channels_first = False try: import torch if isinstance(image, torch.Tensor): t = image.detach().cpu() if t.dtype in (torch.bfloat16, torch.float16): t = t.float() image = t.numpy() channels_first = True except ImportError: # pragma: no cover - torch is a dependency pass if isinstance(image, np.ndarray): return _array_plane(image, bgr, channels_first) raise TypeError( f"unsupported image type {type(image).__name__}: pass a file path, bytes, a PIL image, " "a numpy array or a torch tensor" )