fix(multimodal): keep LLaVA image fetch off the CPU-preprocess timeout budget (flaky test_mixed_batch) (#35700)
This commit is contained in:
@@ -3,6 +3,7 @@ import os
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
from transformers.models.auto.processing_auto import (
|
||||
PROCESSOR_MAPPING_NAMES as HF_MAPPING_NAMES,
|
||||
)
|
||||
@@ -27,7 +28,7 @@ from sglang.srt.multimodal.mm_utils import (
|
||||
process_anyres_image,
|
||||
)
|
||||
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
||||
from sglang.srt.utils import ImageData, load_image, logger
|
||||
from sglang.srt.utils import ImageData, get_image_bytes, load_image, logger
|
||||
from sglang.utils import get_exception_traceback
|
||||
|
||||
|
||||
@@ -44,21 +45,24 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
||||
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def _process_single_image_task(
|
||||
image_data: Union[str, bytes, ImageData],
|
||||
def _preprocess_image_task(
|
||||
image_input,
|
||||
image_hash,
|
||||
image_aspect_ratio: Optional[str] = None,
|
||||
image_grid_pinpoints: Optional[str] = None,
|
||||
processor=None,
|
||||
):
|
||||
|
||||
# CPU-bound decode + preprocessing. `image_input` is either the raw bytes
|
||||
# of a remote image (already fetched off the cpu pool) or a local/inline
|
||||
# input load_image can resolve without network. Either way decode happens
|
||||
# here, parallel across cpu workers, and only compact data crosses the
|
||||
# process boundary (never a decoded image).
|
||||
image_processor = processor.image_processor
|
||||
|
||||
try:
|
||||
url = image_data.url if isinstance(image_data, ImageData) else image_data
|
||||
image, image_size = load_image(url, False)
|
||||
image, image_size = load_image(image_input, False)
|
||||
if image_size is not None:
|
||||
# It is a video with multiple images
|
||||
image_hash = hash(url)
|
||||
pixel_values = image_processor(image)["pixel_values"]
|
||||
for i in range(len(pixel_values)):
|
||||
pixel_values[i] = ensure_numpy(pixel_values[i]).astype(np.float16)
|
||||
@@ -66,7 +70,6 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
||||
return pixel_values, image_hash, image_size
|
||||
else:
|
||||
# It is an image
|
||||
image_hash = hash(url)
|
||||
if image_aspect_ratio == "pad":
|
||||
image = expand2square(
|
||||
image,
|
||||
@@ -93,18 +96,52 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
||||
except Exception:
|
||||
logger.error("Exception in TokenizerManager:\n" + get_exception_traceback())
|
||||
|
||||
async def _fetch_remote_image_bytes(self, url):
|
||||
# Fetch a remote image's compressed bytes in the io thread pool, retrying
|
||||
# only transient network failures. Each attempt is bounded by
|
||||
# REQUEST_TIMEOUT inside download_remote_media, so total time is bounded.
|
||||
loop = asyncio.get_running_loop()
|
||||
max_retries = max(0, int(os.environ.get("SGLANG_MM_LOAD_MAX_RETRIES", "2")))
|
||||
delay = 0.5
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
return await loop.run_in_executor(
|
||||
self.io_executor, get_image_bytes, url
|
||||
)
|
||||
except (
|
||||
requests.exceptions.Timeout,
|
||||
requests.exceptions.ConnectionError,
|
||||
):
|
||||
if attempt >= max_retries:
|
||||
raise
|
||||
await asyncio.sleep(delay)
|
||||
delay *= 2
|
||||
|
||||
async def _process_single_image(
|
||||
self,
|
||||
image_data: Union[bytes, str, ImageData],
|
||||
aspect_ratio: str,
|
||||
grid_pinpoints: str,
|
||||
):
|
||||
url = image_data.url if isinstance(image_data, ImageData) else image_data
|
||||
image_hash = hash(url) if isinstance(url, (str, bytes)) else None
|
||||
|
||||
# Only remote URLs hit the network, and that fetch is what used to sit
|
||||
# inside the cpu-preprocess timeout budget. Pull it into the io pool so a
|
||||
# slow fetch can't starve cpu workers or trip the CPU timeout. Local and
|
||||
# inline inputs (file://, data:, base64, bytes) decode instantly in the
|
||||
# worker, so pass them through unchanged.
|
||||
image_input = url
|
||||
if isinstance(url, str) and url.startswith(("http://", "https://")):
|
||||
image_input = await self._fetch_remote_image_bytes(url)
|
||||
|
||||
if self.cpu_executor is not None:
|
||||
loop = asyncio.get_running_loop()
|
||||
fut = loop.run_in_executor(
|
||||
self.cpu_executor,
|
||||
LlavaImageProcessor._process_single_image_task,
|
||||
image_data,
|
||||
LlavaImageProcessor._preprocess_image_task,
|
||||
image_input,
|
||||
image_hash,
|
||||
aspect_ratio,
|
||||
grid_pinpoints,
|
||||
self._processor,
|
||||
@@ -112,11 +149,12 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
||||
timeout = int(os.environ.get("REQUEST_TIMEOUT", "10"))
|
||||
return await asyncio.wait_for(fut, timeout=timeout)
|
||||
else:
|
||||
return self._process_single_image_task(
|
||||
image_data,
|
||||
return LlavaImageProcessor._preprocess_image_task(
|
||||
image_input,
|
||||
image_hash,
|
||||
aspect_ratio,
|
||||
grid_pinpoints,
|
||||
self._processor.image_processor,
|
||||
self._processor,
|
||||
)
|
||||
|
||||
def _process_precomputed_image_data(self, image_data: List[Dict]) -> Dict:
|
||||
|
||||
Reference in New Issue
Block a user