pcuenq/nvidia-nano-clone
016
1from typing import List, Optional, Union, Any, Dict2 3from PIL import Image4import torch5from transformers.image_processing_base import BatchFeature6from transformers.image_processing_utils_fast import BaseImageProcessorFast, divide_to_patches7from transformers.image_utils import (make_list_of_images, get_image_size,8 get_image_type, ImageInput, ImageType, ChannelDimension)9from transformers.utils import TensorType10import torchvision.transforms as T11 12 13 14class NemotronNanoVLV2ImageProcessor(BaseImageProcessorFast):15 model_input_names = ["pixel_values"]16 17 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, **kwargs):18 super().__init__(**kwargs)19 self.image_size = image_size20 self.max_num_tiles = max_num_tiles21 self.use_thumbnail = use_thumbnail22 self.norm_mean = norm_mean23 self.norm_std = norm_std24 self.do_rescale = do_rescale25 self.num_image_token = int((image_size // patch_size) ** 2 * (downsample_ratio ** 2))26 27 def _process_image(28 self,29 image: ImageInput,30 **kwargs,31 ) -> torch.Tensor:32 image_type = get_image_type(image)33 if image_type == ImageType.PIL:34 if image.mode != 'RGB':35 image = image.convert('RGB')36 image = T.ToTensor()(image)37 return image38 39 def _preprocess(40 self,41 images: List[torch.Tensor],42 image_size: int = None,43 max_num_tiles: int = None,44 use_thumbnail: bool = None,45 do_rescale: bool = None,46 return_tensors: Optional[Union[str, TensorType]] = None,47 **kwargs,48 ) -> List[torch.Tensor]:49 image_size = image_size if image_size is not None else self.image_size50 max_num_tiles = max_num_tiles if max_num_tiles is not None else self.max_num_tiles51 use_thumbnail = use_thumbnail if use_thumbnail is not None else self.use_thumbnail52 do_rescale = do_rescale if do_rescale is not None else self.do_rescale53 54 images = make_list_of_images(images)55 56 all_patches = []57 num_patches = []58 for image in images:59 patches = dynamic_preprocess(image, image_size, max_num_tiles, use_thumbnail)60 all_patches.extend(patches)61 num_patches.append(len(patches))62 63 pixel_values = torch.stack(all_patches, dim=0)64 norm_mean = torch.Tensor(self.norm_mean).view(1, 3, 1, 1)65 norm_std = torch.Tensor(self.norm_std).view(1, 3, 1, 1)66 pixel_values = (pixel_values - norm_mean) / norm_std67 return BatchFeature(data={"pixel_values": pixel_values, "num_patches": num_patches}, tensor_type=return_tensors)68 69 70def get_internvl_target_ratios(71 min_num: int,72 max_num: int,73) -> list[tuple[int, int]]:74 target_ratios = {(i, j)75 for n in range(min_num, max_num + 1)76 for i in range(1, n + 1)77 for j in range(1, n + 1) if min_num <= i * j <= max_num}78 return sorted(target_ratios, key=lambda x: x[0] * x[1])79 80 81# From https://github.com/OpenGVLab/InternVL/blob/c62fa4f7c850165d7386bdc48ac6bc5a6fab0864/internvl_chat/internvl/train/dataset.py#L68582# Copyright (c) 2023 OpenGVLab.83def find_closest_aspect_ratio(84 aspect_ratio: float,85 target_ratios: list[tuple[int, int]],86 width: int,87 height: int,88 image_size: int,89) -> tuple[int, int]:90 best_ratio_diff = float("inf")91 best_ratio = (1, 1)92 area = width * height93 for ratio in target_ratios:94 target_aspect_ratio = ratio[0] / ratio[1]95 ratio_diff = abs(aspect_ratio - target_aspect_ratio)96 if ratio_diff < best_ratio_diff:97 best_ratio_diff = ratio_diff98 best_ratio = ratio99 elif ratio_diff == best_ratio_diff:100 if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:101 best_ratio = ratio102 return best_ratio103 104 105def calculate_targets(106 orig_width: int,107 orig_height: int,108 target_ratios: list[tuple[int, int]],109 image_size: int,110) -> tuple[int, int, int]:111 aspect_ratio = orig_width / orig_height112 113 # find the closest aspect ratio to the target114 target_aspect_ratio = find_closest_aspect_ratio(115 aspect_ratio,116 target_ratios,117 width=orig_width,118 height=orig_height,119 image_size=image_size,120 )121 122 # calculate the target width and height123 target_width = image_size * target_aspect_ratio[0]124 target_height = image_size * target_aspect_ratio[1]125 blocks = target_aspect_ratio[0] * target_aspect_ratio[1]126 127 return blocks, target_width, target_height128 129 130def dynamic_preprocess(image, image_size=512, max_num_tiles=12, use_thumbnail=True):131 orig_height, orig_width = get_image_size(image, channel_dim=ChannelDimension.FIRST)132 target_ratios = get_internvl_target_ratios(1, max_num_tiles)133 134 blocks, target_width, target_height = calculate_targets(135 orig_width,136 orig_height,137 target_ratios,138 image_size139 )140 # resize the image141 resized_img = T.Resize((target_height, target_width), interpolation=T.InterpolationMode.BICUBIC)(image)142 patches = divide_to_patches(resized_img, image_size)143 assert len(patches) == blocks144 if use_thumbnail and len(patches) != 1:145 thumbnail_img = T.Resize((image_size, image_size), interpolation=T.InterpolationMode.BICUBIC)(image)146 patches.append(thumbnail_img)147 148 return patches149 