diff --git a/python/sglang/srt/multimodal/processors/llava.py b/python/sglang/srt/multimodal/processors/llava.py index bbfd41016..ce0ac71f8 100644 --- a/python/sglang/srt/multimodal/processors/llava.py +++ b/python/sglang/srt/multimodal/processors/llava.py @@ -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: