File size: 3,214 Bytes
7325252
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Image preprocessing utilities for HPD-Parsing (transformers path).

Mirrors vLLM's InternVL dynamic tiling path with ``MAX_PATCHES_WITH_RESIZE=true``:
resize to the closest aspect ratio in a ``(min_num, max_num)`` grid, split into
``448x448`` tiles, and optionally append a thumbnail.
"""

import torch
import torchvision.transforms as T
from torchvision.transforms.functional import InterpolationMode
from PIL import Image

IMAGENET_MEAN, IMAGENET_STD = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)

IMAGE_SIZE = 448
MIN_DYNAMIC_PATCH = 1
MAX_DYNAMIC_PATCH = 24
USE_THUMBNAIL = True


def build_transform(input_size=IMAGE_SIZE):
    return T.Compose([
        T.Lambda(lambda img: img.convert("RGB")),
        T.Resize((input_size, input_size), interpolation=InterpolationMode.BICUBIC),
        T.ToTensor(),
        T.Normalize(IMAGENET_MEAN, IMAGENET_STD),
    ])


def get_target_ratios(min_num, max_num):
    ratios = {(i, j)
              for n in range(min_num, max_num + 1)
              for i in range(1, n + 1) for j in range(1, n + 1)
              if min_num <= i * j <= max_num}
    return sorted(ratios, key=lambda x: x[0] * x[1])


def find_closest_aspect_ratio_optim(aspect_ratio, target_ratios, width, height,
                                    image_size, top_k=3, ar_threshold=0.2):
    area = width * height
    candidates = []
    for ratio in target_ratios:
        ar_diff = abs(aspect_ratio - ratio[0] / ratio[1])
        if ar_threshold is not None and ar_diff > ar_threshold:
            continue
        area_diff = abs(area - image_size * image_size * ratio[0] * ratio[1])
        candidates.append((ratio, area_diff, ar_diff))
    if not candidates:  # fall back to no aspect-ratio filtering
        for ratio in target_ratios:
            ar_diff = abs(aspect_ratio - ratio[0] / ratio[1])
            area_diff = abs(area - image_size * image_size * ratio[0] * ratio[1])
            candidates.append((ratio, area_diff, ar_diff))
    candidates.sort(key=lambda x: x[1])
    top = candidates[:top_k]
    top.sort(key=lambda x: x[2])
    return top[0][0]


def dynamic_preprocess(image, target_ratios, image_size=IMAGE_SIZE, use_thumbnail=USE_THUMBNAIL):
    w, h = image.size
    ratio = find_closest_aspect_ratio_optim(w / h, target_ratios, w, h, image_size)
    tw, th = image_size * ratio[0], image_size * ratio[1]
    blocks = ratio[0] * ratio[1]
    resized = image.resize((tw, th))
    cols = tw // image_size
    tiles = []
    for i in range(blocks):
        box = ((i % cols) * image_size, (i // cols) * image_size,
               ((i % cols) + 1) * image_size, ((i // cols) + 1) * image_size)
        tiles.append(resized.crop(box))
    if use_thumbnail and blocks != 1:
        tiles.append(image.resize((image_size, image_size)))
    return tiles


def load_image(path):
    image = Image.open(path).convert("RGB")
    min_num, max_num = MIN_DYNAMIC_PATCH, MAX_DYNAMIC_PATCH
    if USE_THUMBNAIL and max_num != 1:
        max_num += 1
    target_ratios = get_target_ratios(min_num, max_num)
    transform = build_transform(IMAGE_SIZE)
    tiles = dynamic_preprocess(image, target_ratios, IMAGE_SIZE, USE_THUMBNAIL)
    return torch.stack([transform(t) for t in tiles])