File size: 4,775 Bytes
6ffd3f8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
# 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"
    )