image.py 2.96 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
12
13
14
15
16
17
18
19
20
21
22
23


def rescale_image_size(image: Image.Image,
                       size_factor: float,
                       transpose: int = -1) -> Image.Image:
    """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
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
# TODO: Support customizable background color to fill in.
def rgba_to_rgb(
    image: Image.Image, background_color=(255, 255, 255)) -> Image.Image:
    """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)


45
46
47
48
49
50
51
52
53
54
class ImageMediaIO(MediaIO[Image.Image]):

    def __init__(self, *, image_mode: str = "RGB") -> None:
        super().__init__()

        self.image_mode = image_mode

    def load_bytes(self, data: bytes) -> Image.Image:
        image = Image.open(BytesIO(data))
        image.load()
55
        return convert_image_mode(image, self.image_mode)
56
57

    def load_base64(self, media_type: str, data: str) -> Image.Image:
58
        return self.load_bytes(pybase64.b64decode(data, validate=True))
59
60
61
62

    def load_file(self, filepath: Path) -> Image.Image:
        image = Image.open(filepath)
        image.load()
63
        return convert_image_mode(image, self.image_mode)
64
65
66
67
68
69
70
71
72
73

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

        with BytesIO() as buffer:
74
            image = convert_image_mode(image, self.image_mode)
75
76
77
            image.save(buffer, image_format)
            data = buffer.getvalue()

78
        return pybase64.b64encode(data).decode('utf-8')
79
80
81
82
83
84
85
86
87
88
89
90


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:
91
        return self.load_bytes(pybase64.b64decode(data, validate=True))
92
93

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

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