image.py 4.53 KB
Newer Older
1
# SPDX-License-Identifier: Apache-2.0
2
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3

4
5
from io import BytesIO
from pathlib import Path
6

7
import pybase64
8
9
10
import torch
from PIL import Image

11
from .base import MediaIO, MediaWithBytes
12
13


14
15
16
def rescale_image_size(
    image: Image.Image, size_factor: float, transpose: int = -1
) -> Image.Image:
17
18
19
20
21
22
23
    """Rescale the dimensions of an image by a constant factor."""
    new_width = int(image.width * size_factor)
    new_height = int(image.height * size_factor)
    image = image.resize((new_width, new_height))
    if transpose >= 0:
        image = image.transpose(Image.Transpose(transpose))
    return image
24
25


26
def rgba_to_rgb(
27
    image: Image.Image,
28
    background_color: tuple[int, int, int] | list[int] = (255, 255, 255),
29
) -> Image.Image:
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
    """Convert an RGBA image to RGB with filled background color."""
    assert image.mode == "RGBA"
    converted = Image.new("RGB", image.size, background_color)
    converted.paste(image, mask=image.split()[3])  # 3 is the alpha channel
    return converted


def convert_image_mode(image: Image.Image, to_mode: str):
    if image.mode == to_mode:
        return image
    elif image.mode == "RGBA" and to_mode == "RGB":
        return rgba_to_rgb(image)
    else:
        return image.convert(to_mode)


46
class ImageMediaIO(MediaIO[Image.Image]):
47
    def __init__(self, image_mode: str = "RGB", **kwargs) -> None:
48
49
50
        super().__init__()

        self.image_mode = image_mode
51
52
53
54
55
56
        # `kwargs` contains custom arguments from
        # --media-io-kwargs for this modality.
        # They can be passed to the underlying
        # media loaders (e.g. custom implementations)
        # for flexible control.
        self.kwargs = kwargs
57

58
59
        # Extract RGBA background color from kwargs if provided
        # Default to white background for backward compatibility
60
        rgba_bg = kwargs.get("rgba_background_color", (255, 255, 255))
61
62
63
64
65
        # Convert list to tuple for consistency
        if isinstance(rgba_bg, list):
            rgba_bg = tuple(rgba_bg)

        # Validate rgba_background_color format
66
67
68
69
70
        if not (
            isinstance(rgba_bg, tuple)
            and len(rgba_bg) == 3
            and all(isinstance(c, int) and 0 <= c <= 255 for c in rgba_bg)
        ):
71
72
            raise ValueError(
                "rgba_background_color must be a list or tuple of 3 integers "
73
74
                "in the range [0, 255]."
            )
75
76
        self.rgba_background_color = rgba_bg

77
78
79
    def _convert_image_mode(
        self, image: Image.Image | MediaWithBytes[Image.Image]
    ) -> Image.Image:
80
        """Convert image mode with custom background color."""
81
82
        if isinstance(image, MediaWithBytes):
            image = image.media
83
84
85
86
87
88
89
        if image.mode == self.image_mode:
            return image
        elif image.mode == "RGBA" and self.image_mode == "RGB":
            return rgba_to_rgb(image, self.rgba_background_color)
        else:
            return convert_image_mode(image, self.image_mode)

90
    def load_bytes(self, data: bytes) -> MediaWithBytes[Image.Image]:
91
        image = Image.open(BytesIO(data))
92
        return MediaWithBytes(self._convert_image_mode(image), data)
93

94
    def load_base64(self, media_type: str, data: str) -> MediaWithBytes[Image.Image]:
95
        return self.load_bytes(pybase64.b64decode(data, validate=True))
96

97
98
99
100
101
    def load_file(self, filepath: Path) -> MediaWithBytes[Image.Image]:
        with open(filepath, "rb") as f:
            data = f.read()
        image = Image.open(BytesIO(data))
        return MediaWithBytes(self._convert_image_mode(image), data)
102
103
104
105
106
107
108
109
110
111

    def encode_base64(
        self,
        media: Image.Image,
        *,
        image_format: str = "JPEG",
    ) -> str:
        image = media

        with BytesIO() as buffer:
112
            image = self._convert_image_mode(image)
113
114
115
            image.save(buffer, image_format)
            data = buffer.getvalue()

116
        return pybase64.b64encode(data).decode("utf-8")
117
118
119
120
121
122
123
124
125
126
127


class ImageEmbeddingMediaIO(MediaIO[torch.Tensor]):
    def __init__(self) -> None:
        super().__init__()

    def load_bytes(self, data: bytes) -> torch.Tensor:
        buffer = BytesIO(data)
        return torch.load(buffer, weights_only=True)

    def load_base64(self, media_type: str, data: str) -> torch.Tensor:
128
        return self.load_bytes(pybase64.b64decode(data, validate=True))
129
130

    def load_file(self, filepath: Path) -> torch.Tensor:
cyyever's avatar
cyyever committed
131
        return torch.load(filepath, weights_only=True)
132
133

    def encode_base64(self, media: torch.Tensor) -> str:
134
        return pybase64.b64encode(media.numpy()).decode("utf-8")