Download code/tt_superpoint/inputs.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 4.78 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/tt_superpoint/inputs.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/tt_superpoint/inputs.py
-
curl -L -o inputs.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/tt_superpoint/inputs.py
4.78 kB
| # 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" | |
| ) | |