[Feat/WIP] add llava-onevision, with support for (1) siglip encoder, (2) qwen2 decoder (3) openai api compatible server. (#1123)
Co-authored-by: Bo Li <drluodian@gmail.com>
This commit is contained in:
committed by
GitHub
parent
5fafcac008
commit
a5b14ad043
@@ -13,10 +13,25 @@ See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
|
||||
# Source: https://github.com/haotian-liu/LLaVA/blob/main/llava/mm_utils.py
|
||||
# Source: https://github.com/LLaVA-VL/LLaVA-NeXT/blob/main/llava/mm_utils.py
|
||||
"""
|
||||
Utilities for multi-modal models.
|
||||
|
||||
This python file mainly contains utilities that were used in the
|
||||
image processing logic of llava-next including operations such as
|
||||
anyres and anyres_max
|
||||
|
||||
Currently supports the anyres and anyres_max operation for CLIP and
|
||||
SigLip. For more information, you may refer to the paper or the blog
|
||||
|
||||
LLaVA-NeXT : https://llava-vl.github.io/blog/2024-01-30-llava-next/
|
||||
LLaVA-Onevision : https://arxiv.org/pdf/2408.03326
|
||||
|
||||
"""
|
||||
import ast
|
||||
import base64
|
||||
import math
|
||||
import re
|
||||
from io import BytesIO
|
||||
|
||||
import numpy as np
|
||||
@@ -40,10 +55,13 @@ def select_best_resolution(original_size, possible_resolutions):
|
||||
min_wasted_resolution = float("inf")
|
||||
|
||||
for width, height in possible_resolutions:
|
||||
# Calculate the downscaled size to keep the aspect ratio
|
||||
scale = min(width / original_width, height / original_height)
|
||||
downscaled_width, downscaled_height = int(original_width * scale), int(
|
||||
original_height * scale
|
||||
)
|
||||
|
||||
# Calculate effective and wasted resolutions
|
||||
effective_resolution = min(
|
||||
downscaled_width * downscaled_height, original_width * original_height
|
||||
)
|
||||
@@ -129,6 +147,26 @@ def get_anyres_image_grid_shape(image_size, grid_pinpoints, patch_size):
|
||||
Returns:
|
||||
tuple: The shape of the image patch grid in the format (width, height).
|
||||
"""
|
||||
if isinstance(grid_pinpoints, str) and "x" in grid_pinpoints:
|
||||
assert patch_size in [
|
||||
224,
|
||||
336,
|
||||
384,
|
||||
448,
|
||||
512,
|
||||
], "patch_size should be in [224, 336, 384, 448, 512]"
|
||||
# Use regex to extract the range from the input string
|
||||
matches = re.findall(r"\((\d+)x(\d+)\)", grid_pinpoints)
|
||||
range_start = tuple(map(int, matches[0]))
|
||||
range_end = tuple(map(int, matches[-1]))
|
||||
# Generate a matrix of tuples from (range_start[0], range_start[1]) to (range_end[0], range_end[1])
|
||||
grid_pinpoints = [
|
||||
(i, j)
|
||||
for i in range(range_start[0], range_end[0] + 1)
|
||||
for j in range(range_start[1], range_end[1] + 1)
|
||||
]
|
||||
# Multiply all elements by patch_size
|
||||
grid_pinpoints = [[dim * patch_size for dim in pair] for pair in grid_pinpoints]
|
||||
if type(grid_pinpoints) is list:
|
||||
possible_resolutions = grid_pinpoints
|
||||
else:
|
||||
@@ -149,6 +187,31 @@ def process_anyres_image(image, processor, grid_pinpoints):
|
||||
Returns:
|
||||
np.array: An np array containing the processed image patches.
|
||||
"""
|
||||
if isinstance(grid_pinpoints, str) and "x" in grid_pinpoints:
|
||||
try:
|
||||
patch_size = processor.size[0]
|
||||
except Exception as e:
|
||||
patch_size = processor.size["shortest_edge"]
|
||||
assert patch_size in [
|
||||
224,
|
||||
336,
|
||||
384,
|
||||
448,
|
||||
512,
|
||||
], "patch_size should be in [224, 336, 384, 448, 512]"
|
||||
# Use regex to extract the range from the input string
|
||||
matches = re.findall(r"\((\d+)x(\d+)\)", grid_pinpoints)
|
||||
range_start = tuple(map(int, matches[0]))
|
||||
range_end = tuple(map(int, matches[-1]))
|
||||
# Generate a matrix of tuples from (range_start[0], range_start[1]) to (range_end[0], range_end[1])
|
||||
grid_pinpoints = [
|
||||
(i, j)
|
||||
for i in range(range_start[0], range_end[0] + 1)
|
||||
for j in range(range_start[1], range_end[1] + 1)
|
||||
]
|
||||
# Multiply all elements by patch_size
|
||||
grid_pinpoints = [[dim * patch_size for dim in pair] for pair in grid_pinpoints]
|
||||
|
||||
if type(grid_pinpoints) is list:
|
||||
possible_resolutions = grid_pinpoints
|
||||
else:
|
||||
@@ -156,15 +219,24 @@ def process_anyres_image(image, processor, grid_pinpoints):
|
||||
best_resolution = select_best_resolution(image.size, possible_resolutions)
|
||||
image_padded = resize_and_pad_image(image, best_resolution)
|
||||
|
||||
patches = divide_to_patches(image_padded, processor.crop_size["height"])
|
||||
|
||||
image_original_resize = image.resize(
|
||||
(processor.size["shortest_edge"], processor.size["shortest_edge"])
|
||||
# For Siglip processor, only have size but no crop size
|
||||
crop_size = (
|
||||
processor.crop_size["height"]
|
||||
if "crop_size" in processor.__dict__
|
||||
else processor.size["height"]
|
||||
)
|
||||
shortest_edge = (
|
||||
processor.size["shortest_edge"]
|
||||
if "shortest_edge" in processor.size
|
||||
else processor.size["height"]
|
||||
)
|
||||
patches = divide_to_patches(image_padded, crop_size)
|
||||
|
||||
image_original_resize = image.resize((shortest_edge, shortest_edge))
|
||||
|
||||
image_patches = [image_original_resize] + patches
|
||||
image_patches = [
|
||||
processor.preprocess(image_patch)["pixel_values"][0]
|
||||
processor.preprocess(image_patch.convert("RGB"))["pixel_values"][0]
|
||||
for image_patch in image_patches
|
||||
]
|
||||
return np.stack(image_patches, axis=0)
|
||||
@@ -255,7 +327,7 @@ def process_images(images, image_processor, model_cfg):
|
||||
)
|
||||
image = image_processor.preprocess(image)["pixel_values"][0]
|
||||
new_images.append(image)
|
||||
elif image_aspect_ratio == "anyres":
|
||||
elif "anyres" in image_aspect_ratio:
|
||||
for image in images:
|
||||
image = process_anyres_image(
|
||||
image, image_processor, model_cfg.image_grid_pinpoints
|
||||
|
||||
Reference in New Issue
Block a user