| import base64 |
| import functools |
| import io |
| import math |
| from dataclasses import dataclass |
| from typing import Literal, TypedDict |
|
|
| import numpy as np |
| from PIL import Image |
|
|
|
|
| class ImageInput(TypedDict): |
| type: Literal['image'] |
| image: Image.Image |
|
|
|
|
| MediaInput = ImageInput |
|
|
|
|
| @dataclass |
| class TransparentBgConfig: |
| """The config of the transparent background.""" |
|
|
| pattern: Literal["white", "black", "gray", "chessboard"] = "black" |
| """The pattern of the transparent background.""" |
|
|
| chessboard_square_size: int = 16 |
| """The size of the squares in the chessboard background.""" |
|
|
| chessboard_square_on_top_left: bool = True |
| """Whether to start the chessboard with a white square on the top left.""" |
|
|
| chessboard_white_value: int = 255 |
| """The value of the white pixels in the background.""" |
|
|
| chessboard_gray_value: int = 200 |
| """The value of the gray pixels in the background.""" |
|
|
|
|
| @functools.lru_cache(maxsize=256) |
| def _create_chessboard_background( |
| height: int, |
| width: int, |
| square_size: int, |
| square_on_top_left: bool, |
| white_value: int, |
| gray_value: int, |
| ) -> np.ndarray: |
| """Create a chessboard background.""" |
| bg = np.ones((height, width, 3), dtype=np.uint8) * white_value |
| for y in range(0, height, square_size): |
| for x in range(0, width, square_size): |
| if (y // square_size + x // square_size) % 2 == ( |
| 1 if square_on_top_left else 0): |
| bg[y:y + square_size, x:x + square_size] = gray_value |
| return bg |
|
|
|
|
| def fill_transparent_bg_with( |
| image: Image.Image, |
| transparent_bg_config: TransparentBgConfig | None = None, |
| ) -> Image.Image: |
| """Composite a (possibly) transparent image onto a configured background. |
| |
| When ``transparent_bg_config`` is ``None``, the image is simply converted |
| to RGB (preserving the historical behavior). Otherwise the alpha channel |
| is alpha-composited over a background generated according to the config. |
| """ |
| if transparent_bg_config is None: |
| return image.convert("RGB") |
|
|
| if image.mode == "RGB": |
| return image |
|
|
| has_alpha = "A" in image.getbands() or "transparency" in image.info |
| if not has_alpha: |
| return image.convert("RGB") |
|
|
| img = np.array(image.convert("RGBA")) |
| height, width = img.shape[:2] |
| bg_pattern = transparent_bg_config.pattern |
| if bg_pattern == "white": |
| bg = np.full((height, width, 3), 255, dtype=np.uint8) |
| elif bg_pattern == "black": |
| bg = np.zeros((height, width, 3), dtype=np.uint8) |
| elif bg_pattern == "gray": |
| bg = np.full((height, width, 3), 128, dtype=np.uint8) |
| elif bg_pattern == "chessboard": |
| bg = _create_chessboard_background( |
| height, |
| width, |
| transparent_bg_config.chessboard_square_size, |
| transparent_bg_config.chessboard_square_on_top_left, |
| transparent_bg_config.chessboard_white_value, |
| transparent_bg_config.chessboard_gray_value, |
| ) |
| else: |
| raise ValueError(f"Invalid background pattern: {bg_pattern}") |
|
|
| alpha = img[:, :, 3] |
| img_rgb = img[:, :, :3] |
| alpha_normalized = alpha.astype(np.float32) / 255.0 |
| alpha_3d = np.stack([alpha_normalized] * 3, axis=2) |
| result = alpha_3d * img_rgb + (1 - alpha_3d) * bg |
| result = result.astype(np.uint8) |
| return Image.fromarray(result) |
|
|
|
|
| def navit_resize_image( |
| width: int, |
| height: int, |
| patch_size: int, |
| merge_kernel_size: int, |
| in_patch_limit: int, |
| patch_limit_on_one_side: int, |
| fixed_output_tokens: int | None, |
| ): |
| |
| s1 = math.sqrt( |
| in_patch_limit / |
| (max(1.0, width // patch_size) * max(1.0, height // patch_size))) |
| s2 = patch_limit_on_one_side * patch_size / width |
| s3 = patch_limit_on_one_side * patch_size / height |
| scale = min(1.0, s1, s2, s3) |
| new_w, new_h = max(1, int(width * scale)), max(1, int(height * scale)) |
| new_w = min(new_w, patch_limit_on_one_side * patch_size) |
| new_h = min(new_h, patch_limit_on_one_side * patch_size) |
|
|
| |
| factor = merge_kernel_size * patch_size |
|
|
| pad_height = (factor - new_h % factor) % factor |
| pad_width = (factor - new_w % factor) % factor |
|
|
| if fixed_output_tokens is not None: |
| num_tokens = fixed_output_tokens |
| else: |
| |
| token_height = (new_h + pad_height) // factor |
| token_width = (new_w + pad_width) // factor |
|
|
| assert token_height * merge_kernel_size <= patch_limit_on_one_side, ( |
| f"token_height {token_height} * merge_kernel_size {merge_kernel_size} > patch_limit_on_one_side {patch_limit_on_one_side}" |
| ) |
| assert token_width * merge_kernel_size <= patch_limit_on_one_side, ( |
| f"token_width {token_width} * merge_kernel_size {merge_kernel_size} > patch_limit_on_one_side {patch_limit_on_one_side}" |
| ) |
|
|
| num_tokens = token_height * token_width |
| return { |
| "num_tokens": num_tokens, |
| "new_width": new_w, |
| "new_height": new_h, |
| "pad_width": pad_width, |
| "pad_height": pad_height, |
| "sampled_nframes": 1, |
| } |
|
|
|
|
| def _to_pil( |
| data: str | bytes | Image.Image, |
| transparent_bg_config: TransparentBgConfig | None = None, |
| to_rgb: bool = True, |
| ) -> Image.Image: |
| """Load an image and (optionally) composite its transparent background. |
| |
| Args: |
| data: A PIL Image, a base64 ``data:`` URL, a file path, or raw bytes. |
| transparent_bg_config: The config used to fill the transparent |
| background. ``None`` keeps the historical behavior of converting |
| to RGB without compositing. |
| to_rgb: If ``False`` the image is returned as-is (the |
| ``transparent_bg_config`` is ignored). The caller is then |
| expected to call :func:`fill_transparent_bg_with` later — e.g. |
| after a resize. |
| """ |
| if isinstance(data, Image.Image): |
| image = data |
| elif isinstance(data, str): |
| if data.startswith("data:"): |
| raw_base64 = data.split(",")[1] |
| image = Image.open(io.BytesIO(base64.b64decode(raw_base64))) |
| else: |
| image = Image.open(data) |
| elif isinstance(data, bytes): |
| image = Image.open(io.BytesIO(data)) |
| else: |
| raise ValueError(f"Unsupported data type: {type(data)}") |
|
|
| if not to_rgb: |
| return image |
|
|
| return fill_transparent_bg_with(image, transparent_bg_config) |
|
|
|
|
| def ensure_media_type( |
| media: MediaInput, |
| transparent_bg_config: TransparentBgConfig | None = None, |
| transparent_bg_fill_stage: Literal["before_resize", |
| "after_resize"] = "before_resize", |
| ) -> MediaInput: |
| if media['type'] == 'image': |
| media['image'] = _to_pil( |
| media['image'], |
| transparent_bg_config=transparent_bg_config, |
| to_rgb=transparent_bg_fill_stage == "before_resize", |
| ) |
| return media |
| else: |
| raise ValueError(f"Unsupported media type: {media['type']}") |
|
|
|
|
| def image_to_np( |
| image: Image.Image, |
| resize_to: tuple[int, int] | None = None, |
| mode: str = "resize", |
| raise_error_for_ill_resize: bool = True, |
| transparent_bg_config: TransparentBgConfig | None = None, |
| transparent_bg_fill_stage: Literal["before_resize", |
| "after_resize"] = "before_resize", |
| ) -> np.ndarray: |
| """Convert an image to a numpy array. |
| |
| Args: |
| content: The image to convert. |
| resize_to: The size to resize the image to. |
| mode: The mode to resize the image to. |
| raise_error_for_ill_resize: Whether to raise an error for ill-sized resize. |
| transparent_bg_config: The config of the transparent background. Only |
| used when ``transparent_bg_fill_stage == "after_resize"`` (the |
| caller is responsible for filling before resize otherwise). |
| transparent_bg_fill_stage: When to composite the transparent |
| background — before or after the resize step. |
| |
| Returns: |
| A numpy array. |
| """ |
| assert isinstance(image, Image.Image), "image must be a PIL Image" |
| if resize_to is not None: |
| if mode == "resize": |
| image = image.resize(resize_to, resample=Image.Resampling.BICUBIC) |
| if transparent_bg_fill_stage == "after_resize": |
| image = fill_transparent_bg_with(image, transparent_bg_config) |
|
|
| elif mode == "rescale_and_pad_to_center": |
| scale = min(resize_to[0] / image.width, |
| resize_to[1] / image.height, 1.0) |
| new_width = round(image.width * scale) |
| new_height = round(image.height * scale) |
| if new_width == 0 or new_height == 0: |
| if raise_error_for_ill_resize: |
| raise ValueError( |
| f"Invalid resize to: {resize_to}, from image size: {image.size}" |
| ) |
| else: |
| return np.zeros((resize_to[1], resize_to[0], 3), |
| dtype=np.uint8) |
|
|
| image = image.resize((new_width, new_height), |
| resample=Image.Resampling.BICUBIC) |
| if transparent_bg_fill_stage == "after_resize": |
| image = fill_transparent_bg_with(image, transparent_bg_config) |
| padding_left = (resize_to[0] - new_width) // 2 |
| padding_right = resize_to[0] - new_width - padding_left |
| padding_top = (resize_to[1] - new_height) // 2 |
| padding_bottom = resize_to[1] - new_height - padding_top |
| image = np.asarray(image) |
| image = np.pad( |
| image, |
| ((padding_top, padding_bottom), (padding_left, padding_right), |
| (0, 0)), |
| mode="constant", |
| constant_values=0, |
| ) |
| assert image.shape == (resize_to[1], resize_to[0], 3) |
|
|
| elif mode == "rescale_and_pad_to_rightbottom": |
| scale = min(resize_to[0] / image.width, |
| resize_to[1] / image.height, 1.0) |
| new_width = round(image.width * scale) |
| new_height = round(image.height * scale) |
| if new_width == 0 or new_height == 0: |
| if raise_error_for_ill_resize: |
| raise ValueError( |
| f"Invalid resize to: {resize_to}, from image size: {image.size}" |
| ) |
| else: |
| return np.zeros((resize_to[1], resize_to[0], 3), |
| dtype=np.uint8) |
|
|
| image = image.resize((new_width, new_height), |
| resample=Image.Resampling.BICUBIC) |
| if transparent_bg_fill_stage == "after_resize": |
| image = fill_transparent_bg_with(image, transparent_bg_config) |
| padding_right = resize_to[0] - new_width |
| padding_bottom = resize_to[1] - new_height |
| image = np.asarray(image) |
| image = np.pad( |
| image, |
| ((0, padding_bottom), (0, padding_right), (0, 0)), |
| mode="constant", |
| constant_values=0, |
| ) |
| assert image.shape == (resize_to[1], resize_to[0], 3) |
|
|
| else: |
| raise ValueError(f"Invalid mode: {mode}") |
|
|
| if isinstance(image, Image.Image): |
| return np.asarray(image) |
| else: |
| return image |
|
|
|
|
| def navit_patchify(pixel_values: np.ndarray, |
| patch_size: int) -> dict[str, np.ndarray]: |
| """Reshape the pixel values to a navit shape. |
| |
| Args: |
| pixel_values: np.ndarray, shape (t, h, w, c) |
| patch_size: int |
| |
| Returns: |
| dict[str, np.ndarray] |
| - patches: np.ndarray, shape (t * h//patch_size * w//patch_size, c, patch_size, patch_size) |
| - grid_thw: np.ndarray, (t, h//patch_size, w//patch_size) |
| """ |
| T, H, W, C = pixel_values.shape |
| assert C == 3, "pixel_values must have 3 channels" |
|
|
| patches = pixel_values.reshape(T, H // patch_size, patch_size, |
| W // patch_size, patch_size, C) |
| |
| patches = patches.transpose(0, 1, 3, 5, 2, 4) |
| patches = patches.reshape(-1, C, patch_size, patch_size) |
| grid_thw = np.array([T, H // patch_size, W // patch_size]) |
| return {"pixel_values": patches, "grid_thw": grid_thw} |
|
|
|
|
| def normalize(x: np.ndarray, |
| mean, |
| std_inv, |
| pixels_dtype: np.dtype = np.float32) -> np.ndarray: |
| """Normalize the image. |
| |
| Args: |
| x: The image to normalize. The shape is (..., 3). The dtype is uint8. The range is [0, 255]. |
| mean: The mean of the image. |
| std_inv: The inverse of the std of the image. |
| pixels_dtype: The dtype of the image. |
| Returns: |
| The normalized image. The shape is (..., 3). The dtype is determined by the pixels_dtype. |
| """ |
| x = (x / 255.0).astype(pixels_dtype) |
| x -= mean |
| x *= std_inv |
| return x |
|
|
|
|
| def _to_tensor(data, **kwargs): |
| import torch |
|
|
| if isinstance(data, np.ndarray): |
| return torch.from_numpy(data).to(**kwargs) |
| elif isinstance(data, torch.Tensor): |
| return data.to(**kwargs) |
| elif isinstance(data, list): |
| return [_to_tensor(item, **kwargs) for item in data] |
| elif isinstance(data, tuple): |
| return tuple(_to_tensor(item, **kwargs) for item in data) |
| elif isinstance(data, dict): |
| return {k: _to_tensor(v, **kwargs) for k, v in data.items()} |
| elif data is None: |
| return None |
| else: |
| raise ValueError(f"Unsupported data type: {type(data)}") |
|
|