[vlm] fix: recover multimodal decode and processor failures (#36983)
This commit is contained in:
+5
-2
@@ -72,8 +72,11 @@ class MediaSnapshot:
|
|||||||
|
|
||||||
|
|
||||||
def _snapshot_pil(image: Image.Image) -> MediaSnapshot:
|
def _snapshot_pil(image: Image.Image) -> MediaSnapshot:
|
||||||
snapshot = image.copy()
|
try:
|
||||||
snapshot.load()
|
snapshot = image.copy()
|
||||||
|
snapshot.load()
|
||||||
|
except OSError as e:
|
||||||
|
raise ValueError(f"Could not decode image: {e}") from e
|
||||||
payload = snapshot.tobytes()
|
payload = snapshot.tobytes()
|
||||||
palette = snapshot.palette.tobytes() if snapshot.palette is not None else b""
|
palette = snapshot.palette.tobytes() if snapshot.palette is not None else b""
|
||||||
palette_mode = (
|
palette_mode = (
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import dataclasses
|
|||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
import threading
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import (
|
from typing import (
|
||||||
@@ -367,13 +368,8 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.mm_processor_worker_num,
|
self.mm_processor_worker_num,
|
||||||
"auto" if requested_mm_processor_worker_num == 0 else "explicit",
|
"auto" if requested_mm_processor_worker_num == 0 else "explicit",
|
||||||
)
|
)
|
||||||
cpu_worker_start_method = (
|
self._cpu_executor_lock = threading.Lock()
|
||||||
"spawn" if self.mm_feature_transport == "cuda_vmm" else "fork"
|
self.cpu_executor = self._create_cpu_executor()
|
||||||
)
|
|
||||||
self.cpu_executor = concurrent.futures.ProcessPoolExecutor(
|
|
||||||
mp_context=mp.get_context(cpu_worker_start_method),
|
|
||||||
max_workers=int(os.environ.get("SGLANG_CPU_WORKERS", os.cpu_count())),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Mapping from attribute names to modality types
|
# Mapping from attribute names to modality types
|
||||||
self.ATTR_NAME_TO_MODALITY = {
|
self.ATTR_NAME_TO_MODALITY = {
|
||||||
@@ -493,6 +489,41 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
if self.mm_processor_executor is not None:
|
if self.mm_processor_executor is not None:
|
||||||
self.mm_processor_executor.shutdown()
|
self.mm_processor_executor.shutdown()
|
||||||
|
|
||||||
|
def _create_cpu_executor(self) -> concurrent.futures.ProcessPoolExecutor:
|
||||||
|
start_method = "spawn" if self.mm_feature_transport == "cuda_vmm" else "fork"
|
||||||
|
return concurrent.futures.ProcessPoolExecutor(
|
||||||
|
mp_context=mp.get_context(start_method),
|
||||||
|
max_workers=int(os.environ.get("SGLANG_CPU_WORKERS", os.cpu_count())),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _replace_broken_cpu_executor(
|
||||||
|
self, failed_executor: concurrent.futures.ProcessPoolExecutor
|
||||||
|
) -> None:
|
||||||
|
"""Replace a failed preprocess pool once across concurrent requests."""
|
||||||
|
with self._cpu_executor_lock:
|
||||||
|
if self.cpu_executor is not failed_executor:
|
||||||
|
return
|
||||||
|
self.cpu_executor = self._create_cpu_executor()
|
||||||
|
logger.warning("Replaced a broken multimodal CPU preprocess pool")
|
||||||
|
threading.Thread(
|
||||||
|
target=self._shutdown_broken_cpu_executor,
|
||||||
|
args=(failed_executor,),
|
||||||
|
name="sglang-mm-cpu-pool-cleanup",
|
||||||
|
daemon=True,
|
||||||
|
).start()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _shutdown_broken_cpu_executor(
|
||||||
|
failed_executor: concurrent.futures.ProcessPoolExecutor,
|
||||||
|
) -> None:
|
||||||
|
try:
|
||||||
|
failed_executor.shutdown(wait=False, cancel_futures=True)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to shut down a broken multimodal CPU preprocess pool",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
def compute_mrope_positions(self, input_ids, mm_items):
|
def compute_mrope_positions(self, input_ids, mm_items):
|
||||||
"""Compute M-RoPE positions from expanded input_ids and multimodal items.
|
"""Compute M-RoPE positions from expanded input_ids and multimodal items.
|
||||||
|
|
||||||
@@ -863,11 +894,8 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
img, _ = load_image(data, cls.gpu_image_decode)
|
img, _ = load_image(data, cls.gpu_image_decode)
|
||||||
if isinstance(img, torch.Tensor):
|
if isinstance(img, torch.Tensor):
|
||||||
return img # JPEG already decoded on GPU by nvJPEG
|
return img # JPEG already decoded on GPU by nvJPEG
|
||||||
# PIL decodes lazily; do it here in the io worker so the decode
|
|
||||||
# doesn't run later on the event-loop thread.
|
|
||||||
if discard_alpha_channel and img.mode != "RGB":
|
if discard_alpha_channel and img.mode != "RGB":
|
||||||
return img.convert("RGB")
|
return img.convert("RGB")
|
||||||
img.load()
|
|
||||||
return img
|
return img
|
||||||
elif modality == Modality.VIDEO:
|
elif modality == Modality.VIDEO:
|
||||||
return load_video(data, frame_count_limit)
|
return load_video(data, frame_count_limit)
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
|
from concurrent.futures.process import BrokenProcessPool
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import Dict, List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -28,7 +29,13 @@ from sglang.srt.multimodal.mm_utils import (
|
|||||||
process_anyres_image,
|
process_anyres_image,
|
||||||
)
|
)
|
||||||
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
||||||
from sglang.srt.utils import ImageData, get_image_bytes, load_image, logger
|
from sglang.srt.utils import (
|
||||||
|
CLIENT_MEDIA_EXCEPTIONS,
|
||||||
|
ImageData,
|
||||||
|
get_image_bytes,
|
||||||
|
load_image,
|
||||||
|
logger,
|
||||||
|
)
|
||||||
from sglang.utils import get_exception_traceback
|
from sglang.utils import get_exception_traceback
|
||||||
|
|
||||||
|
|
||||||
@@ -93,8 +100,11 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
|||||||
pixel_values = pixel_values.astype(np.float16)
|
pixel_values = pixel_values.astype(np.float16)
|
||||||
|
|
||||||
return pixel_values, image_hash, image.size
|
return pixel_values, image_hash, image.size
|
||||||
|
except CLIENT_MEDIA_EXCEPTIONS as error:
|
||||||
|
raise ValueError(f"Error while processing image: {error}") from error
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.error("Exception in TokenizerManager:\n" + get_exception_traceback())
|
logger.error("Exception in TokenizerManager:\n" + get_exception_traceback())
|
||||||
|
raise
|
||||||
|
|
||||||
async def _fetch_remote_image_bytes(self, url):
|
async def _fetch_remote_image_bytes(self, url):
|
||||||
# Fetch a remote image's compressed bytes in the io thread pool, retrying
|
# Fetch a remote image's compressed bytes in the io thread pool, retrying
|
||||||
@@ -137,17 +147,32 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
|||||||
|
|
||||||
if self.cpu_executor is not None:
|
if self.cpu_executor is not None:
|
||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_running_loop()
|
||||||
fut = loop.run_in_executor(
|
executor = self.cpu_executor
|
||||||
self.cpu_executor,
|
|
||||||
LlavaImageProcessor._preprocess_image_task,
|
|
||||||
image_input,
|
|
||||||
image_hash,
|
|
||||||
aspect_ratio,
|
|
||||||
grid_pinpoints,
|
|
||||||
self._processor,
|
|
||||||
)
|
|
||||||
timeout = int(os.environ.get("REQUEST_TIMEOUT", "10"))
|
timeout = int(os.environ.get("REQUEST_TIMEOUT", "10"))
|
||||||
return await asyncio.wait_for(fut, timeout=timeout)
|
deadline = loop.time() + timeout
|
||||||
|
try:
|
||||||
|
# ProcessPoolExecutor.submit() can itself block after a worker
|
||||||
|
# exits. Keep submission off the request event loop so the
|
||||||
|
# timeout can still replace the failed pool.
|
||||||
|
process_future = await asyncio.wait_for(
|
||||||
|
asyncio.to_thread(
|
||||||
|
executor.submit,
|
||||||
|
LlavaImageProcessor._preprocess_image_task,
|
||||||
|
image_input,
|
||||||
|
image_hash,
|
||||||
|
aspect_ratio,
|
||||||
|
grid_pinpoints,
|
||||||
|
self._processor,
|
||||||
|
),
|
||||||
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
remaining = max(0.0, deadline - loop.time())
|
||||||
|
return await asyncio.wait_for(
|
||||||
|
asyncio.wrap_future(process_future), timeout=remaining
|
||||||
|
)
|
||||||
|
except (BrokenProcessPool, asyncio.TimeoutError):
|
||||||
|
self._replace_broken_cpu_executor(executor)
|
||||||
|
raise
|
||||||
else:
|
else:
|
||||||
return LlavaImageProcessor._preprocess_image_task(
|
return LlavaImageProcessor._preprocess_image_task(
|
||||||
image_input,
|
image_input,
|
||||||
|
|||||||
@@ -413,9 +413,16 @@ class MossVLImageProcessor(SGLangBaseProcessor):
|
|||||||
def _write_video_bytes_to_tempfile(
|
def _write_video_bytes_to_tempfile(
|
||||||
self, video_bytes: bytes, suffix: str = ".mp4"
|
self, video_bytes: bytes, suffix: str = ".mp4"
|
||||||
) -> str:
|
) -> str:
|
||||||
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
|
temp_path = None
|
||||||
f.write(video_bytes)
|
try:
|
||||||
return f.name
|
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
|
||||||
|
temp_path = f.name
|
||||||
|
f.write(video_bytes)
|
||||||
|
return temp_path
|
||||||
|
except BaseException:
|
||||||
|
if temp_path is not None:
|
||||||
|
self._remove_temp_video_paths([temp_path])
|
||||||
|
raise
|
||||||
|
|
||||||
def _normalize_video_string(self, value: str) -> Tuple[str, Optional[str]]:
|
def _normalize_video_string(self, value: str) -> Tuple[str, Optional[str]]:
|
||||||
if value.startswith("file://"):
|
if value.startswith("file://"):
|
||||||
@@ -428,9 +435,8 @@ class MossVLImageProcessor(SGLangBaseProcessor):
|
|||||||
timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
|
timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
|
||||||
content = download_remote_media(value, timeout=timeout)
|
content = download_remote_media(value, timeout=timeout)
|
||||||
suffix = os.path.splitext(urlparse(value).path)[1] or ".mp4"
|
suffix = os.path.splitext(urlparse(value).path)[1] or ".mp4"
|
||||||
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
|
temp_path = self._write_video_bytes_to_tempfile(content, suffix=suffix)
|
||||||
f.write(content)
|
return temp_path, temp_path
|
||||||
return f.name, f.name
|
|
||||||
|
|
||||||
if value.startswith("data:"):
|
if value.startswith("data:"):
|
||||||
header, encoded = value.split(",", 1)
|
header, encoded = value.split(",", 1)
|
||||||
@@ -483,15 +489,46 @@ class MossVLImageProcessor(SGLangBaseProcessor):
|
|||||||
)
|
)
|
||||||
for v in video_data
|
for v in video_data
|
||||||
]
|
]
|
||||||
results = await asyncio.gather(*futures)
|
gather_task = asyncio.gather(*futures, return_exceptions=True)
|
||||||
|
cancelled_error = None
|
||||||
|
try:
|
||||||
|
results = await asyncio.shield(gather_task)
|
||||||
|
except asyncio.CancelledError as error:
|
||||||
|
cancelled_error = error
|
||||||
|
results = await gather_task
|
||||||
|
|
||||||
normalized_inputs: List[Union[str, Dict]] = []
|
normalized_inputs: List[Union[str, Dict]] = []
|
||||||
temp_paths: List[str] = []
|
temp_paths: List[str] = []
|
||||||
for normalized_input, created_paths in results:
|
errors = []
|
||||||
|
for result in results:
|
||||||
|
if isinstance(result, BaseException):
|
||||||
|
errors.append(result)
|
||||||
|
continue
|
||||||
|
normalized_input, created_paths = result
|
||||||
normalized_inputs.append(normalized_input)
|
normalized_inputs.append(normalized_input)
|
||||||
temp_paths.extend(created_paths)
|
temp_paths.extend(created_paths)
|
||||||
|
|
||||||
|
if cancelled_error is not None or errors:
|
||||||
|
self._remove_temp_video_paths(temp_paths)
|
||||||
|
if cancelled_error is not None:
|
||||||
|
raise cancelled_error
|
||||||
|
first_error = errors[0]
|
||||||
|
if len(errors) > 1:
|
||||||
|
first_error.add_note(
|
||||||
|
f"{len(errors) - 1} additional video input(s) failed"
|
||||||
|
)
|
||||||
|
raise first_error
|
||||||
|
|
||||||
return normalized_inputs, temp_paths
|
return normalized_inputs, temp_paths
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _remove_temp_video_paths(temp_paths: List[str]) -> None:
|
||||||
|
for temp_path in temp_paths:
|
||||||
|
try:
|
||||||
|
os.unlink(temp_path)
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
|
||||||
async def process_mm_data_async(
|
async def process_mm_data_async(
|
||||||
self,
|
self,
|
||||||
image_data: List[Union[str, bytes, Dict]],
|
image_data: List[Union[str, bytes, Dict]],
|
||||||
@@ -569,8 +606,4 @@ class MossVLImageProcessor(SGLangBaseProcessor):
|
|||||||
visible_frame_counts=visible_frame_counts,
|
visible_frame_counts=visible_frame_counts,
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
for temp_path in temp_video_paths:
|
self._remove_temp_video_paths(temp_video_paths)
|
||||||
try:
|
|
||||||
os.unlink(temp_path)
|
|
||||||
except FileNotFoundError:
|
|
||||||
pass
|
|
||||||
|
|||||||
@@ -1857,7 +1857,20 @@ def _load_image(
|
|||||||
"Failed to decode JPEG on GPU, falling back to CPU. Error: %s",
|
"Failed to decode JPEG on GPU, falling back to CPU. Error: %s",
|
||||||
e,
|
e,
|
||||||
)
|
)
|
||||||
return Image.open(BytesIO(image_bytes))
|
try:
|
||||||
|
image = Image.open(BytesIO(image_bytes))
|
||||||
|
except OSError as e:
|
||||||
|
raise ValueError(f"Could not decode image: {e}") from e
|
||||||
|
return _fully_load_pil_image(image)
|
||||||
|
|
||||||
|
|
||||||
|
def _fully_load_pil_image(image: Image.Image) -> Image.Image:
|
||||||
|
"""Force PIL's lazy decode while malformed input is still request-local."""
|
||||||
|
try:
|
||||||
|
image.load()
|
||||||
|
except OSError as e:
|
||||||
|
raise ValueError(f"Could not decode image: {e}") from e
|
||||||
|
return image
|
||||||
|
|
||||||
|
|
||||||
def load_image(
|
def load_image(
|
||||||
@@ -1874,7 +1887,7 @@ def load_image(
|
|||||||
image = None
|
image = None
|
||||||
image_size: Optional[tuple[int, int]] = None
|
image_size: Optional[tuple[int, int]] = None
|
||||||
if isinstance(image_file, Image.Image):
|
if isinstance(image_file, Image.Image):
|
||||||
image = image_file
|
image = _fully_load_pil_image(image_file)
|
||||||
image_size = (image.width, image.height)
|
image_size = (image.width, image.height)
|
||||||
elif isinstance(image_file, bytes):
|
elif isinstance(image_file, bytes):
|
||||||
image = _load_image(image_bytes=image_file, gpu_image_decode=gpu_image_decode)
|
image = _load_image(image_bytes=image_file, gpu_image_decode=gpu_image_decode)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import base64
|
||||||
import pickle
|
import pickle
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
@@ -27,7 +28,7 @@ from sglang.srt.disaggregation.encoder.receiver import (
|
|||||||
_encoder_media_item,
|
_encoder_media_item,
|
||||||
_select_mm_processor_prompt,
|
_select_mm_processor_prompt,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.encoder.server import MMEncoder
|
from sglang.srt.disaggregation.encoder.server import BadRequestError, MMEncoder
|
||||||
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
from sglang.srt.managers.tokenizer_manager import (
|
from sglang.srt.managers.tokenizer_manager import (
|
||||||
_reject_missing_dispatched_encoder_embedding,
|
_reject_missing_dispatched_encoder_embedding,
|
||||||
@@ -504,6 +505,17 @@ def test_kimi_k3_epd_selects_matching_jpeg_decode_mode(
|
|||||||
load.assert_called_once_with(b"jpeg", expected_decode_mode)
|
load.assert_called_once_with(b"jpeg", expected_decode_mode)
|
||||||
|
|
||||||
|
|
||||||
|
def test_kimi_k3_epd_rejects_lazy_pil_decode_failure():
|
||||||
|
malformed_png = base64.b64decode(
|
||||||
|
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVQIHWP4z8DwHwAFgAI/ScLJSwAAAABJRU5ErkJggg=="
|
||||||
|
)
|
||||||
|
encoder = _encoder()
|
||||||
|
encoder.preprocessor.use_image_processor_gpu = False
|
||||||
|
|
||||||
|
with pytest.raises(BadRequestError, match="Could not decode image"):
|
||||||
|
encoder.preprocessor._load_single_item(malformed_png, Modality.IMAGE)
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k3_epd_verifies_content_hash_before_decode():
|
def test_kimi_k3_epd_verifies_content_hash_before_decode():
|
||||||
payload = b"jpeg"
|
payload = b"jpeg"
|
||||||
digest = snapshot_media(payload).content_digest
|
digest = snapshot_media(payload).content_digest
|
||||||
|
|||||||
@@ -1,7 +1,10 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from PIL import UnidentifiedImageError
|
||||||
|
|
||||||
from sglang.srt.models.llava import AutoModel, LlavaForConditionalGeneration
|
from sglang.srt.models.llava import AutoModel, LlavaForConditionalGeneration
|
||||||
|
from sglang.srt.multimodal.processors.llava import LlavaImageProcessor
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -87,5 +90,23 @@ class TestLlavaForConditionalGeneration(CustomTestCase):
|
|||||||
self._build_mapping(FakeMapping(ValueError("some other failure")))
|
self._build_mapping(FakeMapping(ValueError("some other failure")))
|
||||||
|
|
||||||
|
|
||||||
|
class TestLlavaImageProcessor(CustomTestCase):
|
||||||
|
@patch("sglang.srt.multimodal.processors.llava.load_image")
|
||||||
|
def test_preprocess_reports_invalid_media_as_client_error(self, mock_load_image):
|
||||||
|
media_error = UnidentifiedImageError("invalid image payload")
|
||||||
|
mock_load_image.side_effect = media_error
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ValueError, "Error while processing image: invalid image payload"
|
||||||
|
) as raised:
|
||||||
|
LlavaImageProcessor._preprocess_image_task(
|
||||||
|
b"invalid",
|
||||||
|
image_hash=1,
|
||||||
|
processor=unittest.mock.Mock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(raised.exception.__cause__, media_error)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,114 @@
|
|||||||
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
from concurrent.futures.process import BrokenProcessPool
|
||||||
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.srt.multimodal.processors.llava import LlavaImageProcessor
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _BrokenExecutor(concurrent.futures.Executor):
|
||||||
|
def submit(self, _fn, /, *_args, **_kwargs):
|
||||||
|
raise BrokenProcessPool("worker exited")
|
||||||
|
|
||||||
|
|
||||||
|
class _BlockingExecutor(concurrent.futures.Executor):
|
||||||
|
def __init__(self):
|
||||||
|
self.release = threading.Event()
|
||||||
|
|
||||||
|
def submit(self, _fn, /, *_args, **_kwargs):
|
||||||
|
self.release.wait(timeout=5)
|
||||||
|
return concurrent.futures.Future()
|
||||||
|
|
||||||
|
|
||||||
|
def _exit_worker_process():
|
||||||
|
os._exit(1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_llava_replaces_broken_pool_without_replaying_request():
|
||||||
|
processor = object.__new__(LlavaImageProcessor)
|
||||||
|
processor.cpu_executor = _BrokenExecutor()
|
||||||
|
processor._processor = Mock()
|
||||||
|
processor._replace_broken_cpu_executor = Mock()
|
||||||
|
|
||||||
|
with pytest.raises(BrokenProcessPool, match="worker exited"):
|
||||||
|
asyncio.run(processor._process_single_image(b"image", "pad", None))
|
||||||
|
|
||||||
|
processor._replace_broken_cpu_executor.assert_called_once_with(
|
||||||
|
processor.cpu_executor
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_llava_times_out_blocked_pool_submission_without_freezing_loop(monkeypatch):
|
||||||
|
monkeypatch.setenv("REQUEST_TIMEOUT", "1")
|
||||||
|
processor = object.__new__(LlavaImageProcessor)
|
||||||
|
processor.cpu_executor = _BlockingExecutor()
|
||||||
|
processor._processor = Mock()
|
||||||
|
processor._replace_broken_cpu_executor = Mock()
|
||||||
|
|
||||||
|
async def run_test():
|
||||||
|
heartbeat = asyncio.Event()
|
||||||
|
|
||||||
|
async def keep_loop_responsive():
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
heartbeat.set()
|
||||||
|
|
||||||
|
heartbeat_task = asyncio.create_task(keep_loop_responsive())
|
||||||
|
try:
|
||||||
|
with pytest.raises(asyncio.TimeoutError):
|
||||||
|
await processor._process_single_image(b"image", "pad", None)
|
||||||
|
assert heartbeat.is_set()
|
||||||
|
finally:
|
||||||
|
processor.cpu_executor.release.set()
|
||||||
|
await heartbeat_task
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
|
processor._replace_broken_cpu_executor.assert_called_once_with(
|
||||||
|
processor.cpu_executor
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_broken_pool_is_replaced_once_for_concurrent_failures():
|
||||||
|
processor = object.__new__(LlavaImageProcessor)
|
||||||
|
failed_executor = Mock()
|
||||||
|
replacement_executor = Mock()
|
||||||
|
processor.cpu_executor = failed_executor
|
||||||
|
processor._cpu_executor_lock = threading.Lock()
|
||||||
|
processor._create_cpu_executor = Mock(return_value=replacement_executor)
|
||||||
|
shutdown_called = threading.Event()
|
||||||
|
failed_executor.shutdown.side_effect = lambda **_kwargs: shutdown_called.set()
|
||||||
|
|
||||||
|
processor._replace_broken_cpu_executor(failed_executor)
|
||||||
|
processor._replace_broken_cpu_executor(failed_executor)
|
||||||
|
|
||||||
|
assert processor.cpu_executor is replacement_executor
|
||||||
|
processor._create_cpu_executor.assert_called_once_with()
|
||||||
|
assert shutdown_called.wait(timeout=5)
|
||||||
|
failed_executor.shutdown.assert_called_once_with(wait=False, cancel_futures=True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_replacement_pool_runs_after_real_worker_exit(monkeypatch):
|
||||||
|
monkeypatch.setenv("SGLANG_CPU_WORKERS", "1")
|
||||||
|
processor = object.__new__(LlavaImageProcessor)
|
||||||
|
processor.mm_feature_transport = "cpu"
|
||||||
|
processor._cpu_executor_lock = threading.Lock()
|
||||||
|
failed_executor = processor._create_cpu_executor()
|
||||||
|
processor.cpu_executor = failed_executor
|
||||||
|
|
||||||
|
try:
|
||||||
|
with pytest.raises(BrokenProcessPool):
|
||||||
|
failed_executor.submit(_exit_worker_process).result(timeout=5)
|
||||||
|
processor._replace_broken_cpu_executor(failed_executor)
|
||||||
|
assert processor.cpu_executor.submit(abs, -1).result(timeout=5) == 1
|
||||||
|
finally:
|
||||||
|
processor.cpu_executor.shutdown(wait=True, cancel_futures=True)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||||
@@ -1,4 +1,7 @@
|
|||||||
|
import asyncio
|
||||||
import re
|
import re
|
||||||
|
import threading
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -81,5 +84,63 @@ def test_moss_vl_accepts_matching_vision_metadata_and_tokens():
|
|||||||
assert rope_deltas.shape == (1,)
|
assert rope_deltas.shape == (1,)
|
||||||
|
|
||||||
|
|
||||||
|
def test_video_normalization_cleans_sibling_temp_file_on_failure(tmp_path):
|
||||||
|
processor = _processor()
|
||||||
|
processor.io_executor = ThreadPoolExecutor(max_workers=2)
|
||||||
|
temp_path = tmp_path / "normalized.mp4"
|
||||||
|
created = threading.Event()
|
||||||
|
|
||||||
|
def normalize(value):
|
||||||
|
if value == "good":
|
||||||
|
temp_path.write_bytes(b"video")
|
||||||
|
created.set()
|
||||||
|
return str(temp_path), [str(temp_path)]
|
||||||
|
assert created.wait(timeout=5)
|
||||||
|
raise ValueError("invalid video")
|
||||||
|
|
||||||
|
processor._normalize_single_video_input = normalize
|
||||||
|
try:
|
||||||
|
with pytest.raises(ValueError, match="invalid video"):
|
||||||
|
asyncio.run(processor._normalize_video_inputs_async(["good", "bad"]))
|
||||||
|
finally:
|
||||||
|
processor.io_executor.shutdown()
|
||||||
|
|
||||||
|
assert not temp_path.exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_video_normalization_waits_for_worker_cleanup_when_cancelled(tmp_path):
|
||||||
|
processor = _processor()
|
||||||
|
processor.io_executor = ThreadPoolExecutor(max_workers=1)
|
||||||
|
temp_path = tmp_path / "cancelled.mp4"
|
||||||
|
created = threading.Event()
|
||||||
|
finish = threading.Event()
|
||||||
|
|
||||||
|
def normalize(_value):
|
||||||
|
temp_path.write_bytes(b"video")
|
||||||
|
created.set()
|
||||||
|
assert finish.wait(timeout=5)
|
||||||
|
return str(temp_path), [str(temp_path)]
|
||||||
|
|
||||||
|
processor._normalize_single_video_input = normalize
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
task = asyncio.create_task(
|
||||||
|
processor._normalize_video_inputs_async(["cancelled"])
|
||||||
|
)
|
||||||
|
assert await asyncio.to_thread(created.wait, 5)
|
||||||
|
task.cancel()
|
||||||
|
finish.set()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
try:
|
||||||
|
asyncio.run(run())
|
||||||
|
finally:
|
||||||
|
finish.set()
|
||||||
|
processor.io_executor.shutdown()
|
||||||
|
|
||||||
|
assert not temp_path.exists()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
raise SystemExit(pytest.main([__file__, "-v"]))
|
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||||
|
|||||||
@@ -8,11 +8,14 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
|||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
import base64
|
||||||
import binascii
|
import binascii
|
||||||
|
import io
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
from sglang.srt.managers.schedule_batch import Modality
|
from sglang.srt.managers.schedule_batch import Modality
|
||||||
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
||||||
@@ -62,6 +65,18 @@ class TestBadInputIsClientError(CustomTestCase):
|
|||||||
# PIL raises UnidentifiedImageError, an OSError -- not a ValueError.
|
# PIL raises UnidentifiedImageError, an OSError -- not a ValueError.
|
||||||
self._assert_client_error(b"definitely not an image", Modality.IMAGE)
|
self._assert_client_error(b"definitely not an image", Modality.IMAGE)
|
||||||
|
|
||||||
|
def test_lazy_pil_decode_failure(self):
|
||||||
|
malformed_png = base64.b64decode(
|
||||||
|
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVQIHWP4z8DwHwAFgAI/ScLJSwAAAABJRU5ErkJggg=="
|
||||||
|
)
|
||||||
|
for media in (malformed_png, Image.open(io.BytesIO(malformed_png))):
|
||||||
|
with self.subTest(media_type=type(media).__name__):
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ValueError, "Could not decode image"
|
||||||
|
) as ctx:
|
||||||
|
_StubProcessor._load_single_item(media, Modality.IMAGE)
|
||||||
|
self.assertIsInstance(ctx.exception.__cause__.__cause__, OSError)
|
||||||
|
|
||||||
def test_undecodable_audio_bytes(self):
|
def test_undecodable_audio_bytes(self):
|
||||||
# soundfile raises LibsndfileError, a RuntimeError -- not a ValueError.
|
# soundfile raises LibsndfileError, a RuntimeError -- not a ValueError.
|
||||||
self._assert_client_error(b"definitely not audio", Modality.AUDIO)
|
self._assert_client_error(b"definitely not audio", Modality.AUDIO)
|
||||||
@@ -91,6 +106,14 @@ class TestServerFaultStaysServerError(CustomTestCase):
|
|||||||
def test_decoder_oom(self):
|
def test_decoder_oom(self):
|
||||||
self._assert_server_error(MemoryError("out of memory"))
|
self._assert_server_error(MemoryError("out of memory"))
|
||||||
|
|
||||||
|
def test_image_source_os_error(self):
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.utils.common.get_image_bytes",
|
||||||
|
side_effect=OSError("too many open files"),
|
||||||
|
):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "too many open files"):
|
||||||
|
_StubProcessor._load_single_item("file:///image.png", Modality.IMAGE)
|
||||||
|
|
||||||
|
|
||||||
class TestClientMediaExceptions(CustomTestCase):
|
class TestClientMediaExceptions(CustomTestCase):
|
||||||
def test_tuple_covers_the_documented_families(self):
|
def test_tuple_covers_the_documented_families(self):
|
||||||
|
|||||||
@@ -1,9 +1,13 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
import unittest
|
import unittest
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from dataclasses import dataclass, replace
|
from dataclasses import dataclass, replace
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
from sglang.srt.managers.schedule_batch import Modality
|
from sglang.srt.managers.schedule_batch import Modality
|
||||||
from sglang.srt.multimodal.cache import MultimodalPreprocessCache, snapshot_media
|
from sglang.srt.multimodal.cache import MultimodalPreprocessCache, snapshot_media
|
||||||
from sglang.srt.multimodal.media_artifacts import (
|
from sglang.srt.multimodal.media_artifacts import (
|
||||||
@@ -85,6 +89,22 @@ class _Processor(MediaArtifactCacheMixin):
|
|||||||
|
|
||||||
|
|
||||||
class TestMediaArtifactProcessor(unittest.TestCase):
|
class TestMediaArtifactProcessor(unittest.TestCase):
|
||||||
|
def test_default_image_decoder_rejects_lazy_pil_failure(self):
|
||||||
|
malformed_png = base64.b64decode(
|
||||||
|
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVQIHWP4z8DwHwAFgAI/ScLJSwAAAABJRU5ErkJggg=="
|
||||||
|
)
|
||||||
|
processor = MediaArtifactCacheMixin()
|
||||||
|
processor.gpu_image_decode = False
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "Could not decode image"):
|
||||||
|
processor.decode_media_snapshot(
|
||||||
|
snapshot_media(malformed_png), Modality.IMAGE
|
||||||
|
)
|
||||||
|
|
||||||
|
lazy_image = Image.open(io.BytesIO(malformed_png))
|
||||||
|
with self.assertRaisesRegex(ValueError, "Could not decode image"):
|
||||||
|
snapshot_media(lazy_image)
|
||||||
|
|
||||||
def test_unknown_model_option_is_part_of_artifact_identity(self):
|
def test_unknown_model_option_is_part_of_artifact_identity(self):
|
||||||
processor = _Processor()
|
processor = _Processor()
|
||||||
digest = snapshot_media(b"image").content_digest
|
digest = snapshot_media(b"image").content_digest
|
||||||
|
|||||||
Reference in New Issue
Block a user