changh95's picture
Python API (2026-10-04): pip install -e code/, from_pretrained() + model(...), Python-first quickstart
6ffd3f8 verified
Raw History Blame Contribute Delete
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"
)