File size: 10,095 Bytes
18218f7 | 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 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 | # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import List, Optional, Union, Any, Dict, Tuple
from PIL import Image
import torch
from transformers.image_processing_base import BatchFeature
from transformers.image_processing_utils_fast import BaseImageProcessorFast
from transformers.image_utils import make_list_of_images, get_image_type, ImageInput, ImageType
from transformers.utils import TensorType
import torchvision.transforms as T
import math
class ZDTaichu5_0_ImageProcessor(BaseImageProcessorFast):
model_input_names = ["pixel_values", "image_grid_thw"]
def __init__(self, image_size=512, max_num_tiles=12, use_thumbnail=True, norm_mean=None, norm_std=None, do_rescale=True, patch_size=16, downsample_ratio=0.5, merge_size=1, **kwargs):
super().__init__(**kwargs)
self.image_size = image_size
self.max_num_tiles = max_num_tiles
self.use_thumbnail = use_thumbnail
self.norm_mean = norm_mean
self.norm_std = norm_std
self.do_rescale = do_rescale
self.merge_size = merge_size
self.num_image_token = int((image_size // patch_size) ** 2 * (downsample_ratio ** 2))
def _process_image(
self,
image: ImageInput,
**kwargs,
) -> torch.Tensor:
image_type = get_image_type(image)
if image_type == ImageType.PIL:
if image.mode != 'RGB':
image = image.convert('RGB')
# Keep PIL input through tiling so resize order matches vLLM.
return image
def _preprocess(
self,
images: List[torch.Tensor],
image_size: int = None,
max_num_tiles: int = None,
use_thumbnail: bool = None,
do_rescale: bool = None,
return_tensors: Optional[Union[str, TensorType]] = None,
**kwargs,
) -> List[torch.Tensor]:
image_size = image_size if image_size is not None else self.image_size
max_num_tiles = max_num_tiles if max_num_tiles is not None else self.max_num_tiles
use_thumbnail = use_thumbnail if use_thumbnail is not None else self.use_thumbnail
do_rescale = do_rescale if do_rescale is not None else self.do_rescale
images = make_list_of_images(images)
all_patches = []
num_patches = []
image_grid_thw = []
for image in images:
patches, tile_rows, tile_cols = dynamic_preprocess(image, image_size, max_num_tiles, use_thumbnail)
all_patches.extend(patches)
num_patches.append(len(patches))
image_grid_thw.append([1, tile_rows, tile_cols])
# vLLM converts each already-cropped PIL tile with ToTensor.
pixel_values = torch.stack([T.ToTensor()(patch) for patch in all_patches], dim=0)
norm_mean = torch.Tensor(self.norm_mean).view(1, 3, 1, 1)
norm_std = torch.Tensor(self.norm_std).view(1, 3, 1, 1)
pixel_values = (pixel_values - norm_mean) / norm_std
pixel_values = pixel_values.to(torch.bfloat16)
return BatchFeature(
data={
"pixel_values": pixel_values,
"num_patches": num_patches,
"image_grid_thw": image_grid_thw,
},
tensor_type=return_tensors,
)
def get_internvl_target_ratios(
min_num: int,
max_num: int,
) -> list[tuple[int, int]]:
target_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(target_ratios, key=lambda x: x[0] * x[1])
# From https://github.com/OpenGVLab/InternVL/blob/c62fa4f7c850165d7386bdc48ac6bc5a6fab0864/internvl_chat/internvl/train/dataset.py#L685
# Copyright (c) 2023 OpenGVLab.
def find_closest_aspect_ratio(
aspect_ratio: float,
target_ratios: list[tuple[int, int]],
width: int,
height: int,
image_size: int,
) -> tuple[int, int]:
best_ratio_diff = float("inf")
best_ratio = (1, 1)
area = width * height
for ratio in target_ratios:
target_aspect_ratio = ratio[0] / ratio[1]
ratio_diff = abs(aspect_ratio - target_aspect_ratio)
if ratio_diff < best_ratio_diff:
best_ratio_diff = ratio_diff
best_ratio = ratio
elif ratio_diff == best_ratio_diff:
if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
best_ratio = ratio
return best_ratio
def select_tile_grid(
*,
orig_width: int,
orig_height: int,
image_size: int,
min_num_tiles: int,
max_num_tiles: int,
) -> Tuple[int, int, int, int]:
"""
Choose (rw, rh) tile grid with small-image and aspect-sanity guards.
Returns (num_grid_blocks, target_width, target_height, effective_max_tiles).
num_grid_blocks = rw * rh (does NOT include the optional thumbnail).
Guards:
1. Area cap: don't create more tiles than the source has pixels for.
2. Aspect-sanity: drop candidates whose ratio differs from source by >3x.
"""
# ββ Guard 1: area cap βββββββββββββββββββββββββββββββββββββββββββββββββ
src_pixels = orig_width * orig_height
tile_pixels = image_size * image_size
area_max_tiles = max(1, math.ceil(src_pixels / tile_pixels))
effective_max = min(max_num_tiles, area_max_tiles)
effective_max = max(effective_max, min_num_tiles)
target_ratios = get_internvl_target_ratios(min_num_tiles, effective_max)
# ββ Guard 2: aspect-sanity ββββββββββββββββββββββββββββββββββββββββββββ
src_ar = orig_width / orig_height
filtered = [
(rw, rh) for (rw, rh) in target_ratios
if (1.0 / 3.0) <= (rw / rh) / src_ar <= 3.0
]
# Fall back to unfiltered set for extreme panoramas / long strips where
# no candidate is within 3x β better to pick *something* than error.
if filtered:
target_ratios = filtered
# ββ Pick best ratio βββββββββββββββββββββββββββββββββββββββββββββββββββ
rw, rh = find_closest_aspect_ratio(
src_ar, target_ratios,
width=orig_width, height=orig_height, image_size=image_size,
)
target_width = image_size * rw
target_height = image_size * rh
num_grid_blocks = rw * rh
return num_grid_blocks, target_width, target_height, effective_max
def count_tiles(
*,
orig_width: int,
orig_height: int,
image_size: int,
min_num_tiles: int,
max_num_tiles: int,
use_thumbnail: bool,
) -> int:
"""
Total number of tiles (grid blocks + optional thumbnail) for this image.
This MUST match what the actual tiling produces, or prompt expansion and
embedding count will diverge.
"""
n_grid, _, _, _ = select_tile_grid(
orig_width=orig_width, orig_height=orig_height,
image_size=image_size,
min_num_tiles=min_num_tiles, max_num_tiles=max_num_tiles,
)
if use_thumbnail and n_grid != 1:
return n_grid + 1
return n_grid
def calculate_targets(
orig_width: int,
orig_height: int,
target_ratios: list[tuple[int, int]],
image_size: int,
) -> tuple[int, int, int]:
aspect_ratio = orig_width / orig_height
# find the closest aspect ratio to the target
target_aspect_ratio = find_closest_aspect_ratio(
aspect_ratio,
target_ratios,
width=orig_width,
height=orig_height,
image_size=image_size,
)
# calculate the target width and height
target_width = image_size * target_aspect_ratio[0]
target_height = image_size * target_aspect_ratio[1]
blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
return blocks, target_width, target_height
def dynamic_preprocess(image, image_size=512, max_num_tiles=12, use_thumbnail=True, min_num_tiles=1):
"""Split a PIL image using vLLM's resize/crop order."""
if isinstance(image, torch.Tensor):
image = T.ToPILImage()(image)
elif not isinstance(image, Image.Image):
image = Image.fromarray(image)
if image.mode != 'RGB':
image = image.convert('RGB')
orig_width, orig_height = image.size
n_grid, target_width, target_height, _ = select_tile_grid(
orig_width=orig_width,
orig_height=orig_height,
image_size=image_size,
min_num_tiles=min_num_tiles,
max_num_tiles=max_num_tiles,
)
# Tile grid dimensions (rows Γ cols of the InternVL tiling)
tile_rows = target_height // image_size
tile_cols = target_width // image_size
resized_img = image.resize((target_width, target_height), Image.BICUBIC)
cols = target_width // image_size
patches = []
for i in range(n_grid):
col = i % cols
row = i // cols
patches.append(
resized_img.crop(
(col * image_size, row * image_size,
(col + 1) * image_size, (row + 1) * image_size)
)
)
assert len(patches) == n_grid
if use_thumbnail and n_grid != 1:
thumbnail = image.resize((image_size, image_size), Image.BICUBIC)
patches.append(thumbnail)
#print(orig_height, orig_width, target_width, target_height, len(patches))
return patches, tile_rows, tile_cols
|