[Fix] Drop deprecated multimodal processor residency state (#33308)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Liangsheng Yin
2026-08-02 20:02:53 -07:00
committed by GitHub
co-authored by Mick
parent 28a2472f95
commit dd6ddc053b
8 changed files with 213 additions and 55 deletions
@@ -337,6 +337,10 @@ class MultimodalDataItem:
def set(self, key: str, value: Any): def set(self, key: str, value: Any):
self.__setitem__(key, value) self.__setitem__(key, value)
def set_hash(self, hash_value: int) -> None:
self.hash = hash_value
self.pad_value = _compute_pad_value(hash_value)
@staticmethod @staticmethod
def is_empty_list(l): def is_empty_list(l):
if l is None: if l is None:
@@ -157,6 +157,23 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _reject_missing_dispatched_encoder_embedding(server_args, request_obj, mm_inputs):
"""Do not silently turn a failed EPD request into local vision work."""
if (
mm_inputs is None
and server_args.language_only
and server_args.encoder_transfer_backend == "zmq_to_tokenizer"
and request_obj.need_wait_for_mm_inputs
):
raise fastapi.HTTPException(
status_code=HTTPStatus.SERVICE_UNAVAILABLE,
detail=(
"The encoder did not return multimodal embeddings. "
"The request was not run locally in language-only mode."
),
)
@lru_cache(maxsize=1) @lru_cache(maxsize=1)
def _ragged_verify_cap_accept() -> bool: def _ragged_verify_cap_accept() -> bool:
# The mode env is fixed at server launch; cache to keep it off the # The mode env is fixed at server launch; cache to keep it off the
@@ -983,6 +1000,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._validate_mm_limits(obj) self._validate_mm_limits(obj)
mm_inputs = None mm_inputs = None
mm_processor_input = (
input_ids
if self.mm_processor.prefer_tokenized_input and input_ids is not None
else (input_text or input_ids)
)
if ( if (
not self.server_args.language_only not self.server_args.language_only
@@ -992,9 +1014,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
mm_inputs = await self.mm_receiver.recv_mm_data( mm_inputs = await self.mm_receiver.recv_mm_data(
request_obj=obj, request_obj=obj,
mm_processor=self.mm_processor, mm_processor=self.mm_processor,
prompt=(input_text or input_ids), prompt=mm_processor_input,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
) )
_reject_missing_dispatched_encoder_embedding(
self.server_args, obj, mm_inputs
)
if mm_inputs is None: if mm_inputs is None:
if self.server_args.language_only: if self.server_args.language_only:
logger.warning( logger.warning(
@@ -1004,7 +1029,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
mm_inputs = await self.mm_processor.process_mm_data_async( mm_inputs = await self.mm_processor.process_mm_data_async(
image_data=obj.image_data, image_data=obj.image_data,
audio_data=obj.audio_data, audio_data=obj.audio_data,
input_text=(input_text or input_ids), input_text=mm_processor_input,
request_obj=obj, request_obj=obj,
max_req_input_len=self.max_req_input_len, max_req_input_len=self.max_req_input_len,
) )
@@ -1019,7 +1044,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
mm_inputs = await self.mm_processor.process_mm_data_async( mm_inputs = await self.mm_processor.process_mm_data_async(
image_data=obj.image_data, image_data=obj.image_data,
audio_data=obj.audio_data, audio_data=obj.audio_data,
input_text=(input_text or input_ids), input_text=mm_processor_input,
request_obj=obj, request_obj=obj,
max_req_input_len=self.max_req_input_len, max_req_input_len=self.max_req_input_len,
) )
@@ -1054,7 +1079,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if not isinstance(item, MultimodalDataItem): if not isinstance(item, MultimodalDataItem):
continue continue
try: try:
item.hash = int(hex_hash, 16) item.set_hash(int(hex_hash, 16))
except (TypeError, ValueError): except (TypeError, ValueError):
logger.warning( logger.warning(
"Ignoring malformed mm_hashes entry %r; " "Ignoring malformed mm_hashes entry %r; "
@@ -44,8 +44,6 @@ _is_cpu = is_cpu()
_is_npu = is_npu() _is_npu = is_npu()
_is_xpu = is_xpu() _is_xpu = is_xpu()
_IPC_POOL_HANDLE_CACHE = envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
@dataclasses.dataclass @dataclasses.dataclass
class BaseMultiModalProcessorOutput: class BaseMultiModalProcessorOutput:
@@ -182,6 +180,8 @@ class MultimodalSpecialTokens:
class BaseMultimodalProcessor(ABC): class BaseMultimodalProcessor(ABC):
models = [] models = []
gpu_image_decode = True # Enable GPU decoding by default gpu_image_decode = True # Enable GPU decoding by default
prefer_tokenized_input = False
precompute_hash_before_cpu_transfer = False
auto_mm_processor_worker_num = 1 auto_mm_processor_worker_num = 1
auto_mm_io_worker_num = 4 auto_mm_io_worker_num = 4
supports_mm_processor_concurrency = False supports_mm_processor_concurrency = False
@@ -193,7 +193,6 @@ class BaseMultimodalProcessor(ABC):
self._processor = _processor self._processor = _processor
self.server_args = server_args self.server_args = server_args
self.transport_mode = transport_mode self.transport_mode = transport_mode
self.keep_mm_feature_on_device = server_args.keep_mm_feature_on_device
configured_mm_feature_transport = getattr( configured_mm_feature_transport = getattr(
server_args, "mm_feature_transport", "cpu" server_args, "mm_feature_transport", "cpu"
) )
@@ -203,6 +202,9 @@ class BaseMultimodalProcessor(ABC):
else "cpu" else "cpu"
) )
self.use_cuda_ipc = self.mm_feature_transport == "cuda_ipc" self.use_cuda_ipc = self.mm_feature_transport == "cuda_ipc"
self.use_ipc_pool_handle_cache = (
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
)
self.disable_fast_image_processor = server_args.disable_fast_image_processor self.disable_fast_image_processor = server_args.disable_fast_image_processor
self.skip_tokenizer_init = server_args.skip_tokenizer_init self.skip_tokenizer_init = server_args.skip_tokenizer_init
@@ -573,16 +575,15 @@ class BaseMultimodalProcessor(ABC):
return_tensors="pt", return_tensors="pt",
**kwargs, **kwargs,
) )
if not self.keep_mm_feature_on_device: # Deferred: the hash is computed on the GPU tensor first, and
# _precompute_hashes_before_cpu_transfer moves it down afterwards.
if not self.use_cuda_ipc and not self.precompute_hash_before_cpu_transfer:
# move feature tensors to cpu # move feature tensors to cpu
for feature_name in self.FEATURE_NAMES: for feature_name in self.FEATURE_NAMES:
if self.use_cuda_ipc: if feature_name in result and isinstance(
pass result[feature_name], torch.Tensor
else: ):
if feature_name in result and isinstance( result[feature_name] = result[feature_name].to("cpu")
result[feature_name], torch.Tensor
):
result[feature_name] = result[feature_name].to("cpu")
return result return result
@@ -1019,13 +1020,17 @@ class BaseMultimodalProcessor(ABC):
for modality, idx, future in futures: for modality, idx, future in futures:
try: try:
result = await asyncio.wrap_future(future) result = await asyncio.wrap_future(future)
except ValueError: except ValueError as e:
logger.exception( logger.info(
"[load_mm_data(simple)] error loading %s data at index=%d", "[load_mm_data(simple)] invalid %s data at index=%d: %s",
modality.name, modality.name,
idx, idx,
e,
) )
raise raise ValueError(
f"An exception occurred while loading {modality.name} data "
f"at index {idx}: {e}"
) from e
except Exception as e: except Exception as e:
logger.exception( logger.exception(
"[load_mm_data(simple)] error loading %s data at index=%d", "[load_mm_data(simple)] error loading %s data at index=%d",
@@ -1167,6 +1172,10 @@ class BaseMultimodalProcessor(ABC):
raise RuntimeError( raise RuntimeError(
f"An exception occurred while loading multimodal data: {e}" f"An exception occurred while loading multimodal data: {e}"
) )
except ValueError as e:
raise ValueError(
f"An exception occurred while loading multimodal data: {e}"
) from e
except Exception as e: except Exception as e:
raise RuntimeError( raise RuntimeError(
f"An exception occurred while loading multimodal data: {e}" f"An exception occurred while loading multimodal data: {e}"
@@ -1349,16 +1358,38 @@ class BaseMultimodalProcessor(ABC):
sync_buffer_meta=sync_flag, sync_buffer_meta=sync_flag,
pool_ipc_handle=( pool_ipc_handle=(
self.cudaipc_mmfeature_pool._pool_ipc_handle self.cudaipc_mmfeature_pool._pool_ipc_handle
if _IPC_POOL_HANDLE_CACHE if self.use_ipc_pool_handle_cache
else None else None
), ),
pool_byte_offset=byte_offset, pool_byte_offset=byte_offset,
pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index, pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index,
) )
if self.keep_mm_feature_on_device:
return tensor
return tensor.cpu() return tensor.cpu()
@staticmethod
def _move_feature_to_cpu(value):
if isinstance(value, torch.Tensor):
return value.cpu()
if isinstance(value, list):
return [BaseMultimodalProcessor._move_feature_to_cpu(v) for v in value]
if isinstance(value, tuple):
return tuple(BaseMultimodalProcessor._move_feature_to_cpu(v) for v in value)
return value
def _precompute_hashes_before_cpu_transfer(
self, mm_items: List[MultimodalDataItem]
) -> None:
if not self.precompute_hash_before_cpu_transfer:
return
for item in mm_items:
item.set_pad_value()
if not self.use_cuda_ipc:
item.feature = self._move_feature_to_cpu(item.feature)
item.precomputed_embeddings = self._move_feature_to_cpu(
item.precomputed_embeddings
)
def resolve_image_token_counts(self, images: List) -> List[int]: def resolve_image_token_counts(self, images: List) -> List[int]:
"""Per-image expanded token counts, computed without re-tokenizing. """Per-image expanded token counts, computed without re-tokenizing.
@@ -1577,14 +1608,10 @@ class BaseMultimodalProcessor(ABC):
): ):
item.set_pad_value() item.set_pad_value()
""" self._precompute_hashes_before_cpu_transfer(all_collected_items)
solution for cuda-ipc memory-leak:
1. memory-pool: each time get a slice from memory-pool and use it as transport-data (with async lock guard)
2. if can not get a slice , transport normal tensor
3. copy tensor in scheduler and release it (use position mark)
4. copy
"""
# Wrap GPU features in the bounded IPC pool; pool misses fall back to a
# plain CPU tensor. The scheduler copies out and releases each slice.
if self.use_cuda_ipc: if self.use_cuda_ipc:
# post-process, prepare for cuda-ipc transfer # post-process, prepare for cuda-ipc transfer
for item in all_collected_items: for item in all_collected_items:
@@ -346,16 +346,13 @@ class Ernie4_5_VLImageProcessor(SGLangBaseProcessor):
if result["pixel_values_videos"].numel() == 0: if result["pixel_values_videos"].numel() == 0:
del result["pixel_values_videos"] del result["pixel_values_videos"]
if not self.keep_mm_feature_on_device: if not self.use_cuda_ipc:
# move feature tensors to cpu # move feature tensors to cpu
for feature_name in self.FEATURE_NAMES: for feature_name in self.FEATURE_NAMES:
if self.use_cuda_ipc: if feature_name in result and isinstance(
pass result[feature_name], torch.Tensor
else: ):
if feature_name in result and isinstance( result[feature_name] = result[feature_name].to("cpu")
result[feature_name], torch.Tensor
):
result[feature_name] = result[feature_name].to("cpu")
return result return result
@@ -416,6 +416,8 @@ class KimiGPUProcessorWrapper:
class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor): class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
models = [KimiK25ForConditionalGeneration] models = [KimiK25ForConditionalGeneration]
gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others
prefer_tokenized_input = True
precompute_hash_before_cpu_transfer = True
def __init__(self, hf_config, server_args, _processor, *args, **kwargs): def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
super().__init__(hf_config, server_args, _processor, *args, **kwargs) super().__init__(hf_config, server_args, _processor, *args, **kwargs)
@@ -70,7 +70,7 @@ class MiDashengLMMultimodalProcessor(BaseMultimodalProcessor):
**kwargs, **kwargs,
) )
if not self.keep_mm_feature_on_device and not self.use_cuda_ipc: if not self.use_cuda_ipc:
for feature_name in ["input_values"]: for feature_name in ["input_values"]:
if feature_name in result: if feature_name in result:
result[feature_name] = result[feature_name].cpu() result[feature_name] = result[feature_name].cpu()
@@ -4,14 +4,17 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-small") register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=1, suite="stage-b-test-1-gpu-small-amd")
class TestMmProcessConfigValidation(unittest.TestCase): class TestMmProcessConfigValidation(CustomTestCase):
"""Server-args validation for mm_process_config.""" """Server-args validation for mm_process_config."""
def _validate_config(self, mm_process_config): def _validate_config(self, mm_process_config):
@@ -63,7 +66,7 @@ class TestMmProcessConfigValidation(unittest.TestCase):
self.assertEqual(args.mm_process_config, config) self.assertEqual(args.mm_process_config, config)
class TestBaseProcessorConfigExtraction(unittest.TestCase): class TestBaseProcessorConfigExtraction(CustomTestCase):
"""Verify BaseMultimodalProcessor.__init__ extracts configs from server_args.""" """Verify BaseMultimodalProcessor.__init__ extracts configs from server_args."""
def _make_processor( def _make_processor(
@@ -159,12 +162,11 @@ class TestBaseProcessorConfigExtraction(unittest.TestCase):
self.assertEqual(proc.mm_io_worker_num, 6) self.assertEqual(proc.mm_io_worker_num, 6)
class TestMultimodalFeatureTransportRuntime(unittest.TestCase): class TestMultimodalFeatureTransportRuntime(CustomTestCase):
@staticmethod @staticmethod
def _server_args(mm_feature_transport): def _server_args(mm_feature_transport):
return SimpleNamespace( return SimpleNamespace(
mm_feature_transport=mm_feature_transport, mm_feature_transport=mm_feature_transport,
keep_mm_feature_on_device=False,
disable_fast_image_processor=False, disable_fast_image_processor=False,
skip_tokenizer_init=False, skip_tokenizer_init=False,
mm_process_config={}, mm_process_config={},
@@ -185,7 +187,7 @@ class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
# transport policy must still resolve from the instance's ServerArgs. # transport policy must still resolve from the instance's ServerArgs.
from sglang.srt.multimodal.processors import base_processor from sglang.srt.multimodal.processors import base_processor
with patch.object( with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True), patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set() base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool: ), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
processor = base_processor.BaseMultimodalProcessor( processor = base_processor.BaseMultimodalProcessor(
@@ -197,12 +199,30 @@ class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
self.assertEqual(processor.mm_feature_transport, "cuda_ipc") self.assertEqual(processor.mm_feature_transport, "cuda_ipc")
self.assertTrue(processor.use_cuda_ipc) self.assertTrue(processor.use_cuda_ipc)
self.assertTrue(processor.use_ipc_pool_handle_cache)
memory_pool.assert_called_once()
def test_cuda_ipc_pool_handle_cache_can_be_disabled(self):
from sglang.srt.multimodal.processors import base_processor
with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(False), patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cuda_ipc"),
_processor=self._processor(),
transport_mode=None,
)
self.assertTrue(processor.use_cuda_ipc)
self.assertFalse(processor.use_ipc_pool_handle_cache)
memory_pool.assert_called_once() memory_pool.assert_called_once()
def test_cpu_transport_does_not_allocate_ipc_pool(self): def test_cpu_transport_does_not_allocate_ipc_pool(self):
from sglang.srt.multimodal.processors import base_processor from sglang.srt.multimodal.processors import base_processor
with patch.object( with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True), patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set() base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool: ), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
processor = base_processor.BaseMultimodalProcessor( processor = base_processor.BaseMultimodalProcessor(
@@ -214,9 +234,51 @@ class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
self.assertEqual(processor.mm_feature_transport, "cpu") self.assertEqual(processor.mm_feature_transport, "cpu")
self.assertFalse(processor.use_cuda_ipc) self.assertFalse(processor.use_cuda_ipc)
self.assertFalse(processor.use_ipc_pool_handle_cache)
memory_pool.assert_not_called() memory_pool.assert_not_called()
class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
@staticmethod
def _processor(enabled):
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
)
with patch.object(
BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
processor = BaseMultimodalProcessor()
processor.precompute_hash_before_cpu_transfer = enabled
processor.use_cuda_ipc = False
return processor
def test_enabled_path_sets_hash_and_pad_value(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
item = MultimodalDataItem(
modality=Modality.IMAGE, feature=torch.arange(8, dtype=torch.float32)
)
self._processor(True)._precompute_hashes_before_cpu_transfer([item])
self.assertIsNotNone(item.hash)
self.assertIsNotNone(item.pad_value)
self.assertTrue(item.feature.is_cpu)
def test_disabled_path_leaves_item_unmodified(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
item = MultimodalDataItem(
modality=Modality.IMAGE, feature=torch.arange(8, dtype=torch.float32)
)
self._processor(False)._precompute_hashes_before_cpu_transfer([item])
self.assertIsNone(item.hash)
self.assertIsNone(item.pad_value)
class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase): class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase):
async def test_dedicated_executor_runs_processor_off_event_loop(self): async def test_dedicated_executor_runs_processor_off_event_loop(self):
from sglang.srt.multimodal.processors.base_processor import ( from sglang.srt.multimodal.processors.base_processor import (
@@ -311,7 +373,7 @@ class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase):
self.assertEqual(deepcopy.call_count, 3) self.assertEqual(deepcopy.call_count, 3)
class TestProcessMmDataKwargs(unittest.TestCase): class TestProcessMmDataKwargs(CustomTestCase):
"""Verify process_mm_data injects per-modality kwargs correctly.""" """Verify process_mm_data injects per-modality kwargs correctly."""
def _make_base_processor(self, mm_process_config): def _make_base_processor(self, mm_process_config):
@@ -324,7 +386,6 @@ class TestProcessMmDataKwargs(unittest.TestCase):
server_args.mm_process_config = mm_process_config server_args.mm_process_config = mm_process_config
server_args.mm_feature_transport = "cpu" server_args.mm_feature_transport = "cpu"
server_args.disable_fast_image_processor = True server_args.disable_fast_image_processor = True
server_args.keep_mm_feature_on_device = True
server_args.skip_tokenizer_init = False server_args.skip_tokenizer_init = False
mock_processor = MagicMock() mock_processor = MagicMock()
@@ -343,7 +404,6 @@ class TestProcessMmDataKwargs(unittest.TestCase):
proc = BaseMultimodalProcessor() proc = BaseMultimodalProcessor()
proc.server_args = server_args proc.server_args = server_args
proc.keep_mm_feature_on_device = server_args.keep_mm_feature_on_device
proc.mm_feature_transport = server_args.mm_feature_transport proc.mm_feature_transport = server_args.mm_feature_transport
proc.use_cuda_ipc = False proc.use_cuda_ipc = False
proc.disable_fast_image_processor = server_args.disable_fast_image_processor proc.disable_fast_image_processor = server_args.disable_fast_image_processor
@@ -452,7 +512,7 @@ class TestProcessMmDataKwargs(unittest.TestCase):
self.assertEqual(audio_kw.get("sample_rate"), 16000) self.assertEqual(audio_kw.get("sample_rate"), 16000)
class TestOverrideProcessorsConfigInjection(unittest.TestCase): class TestOverrideProcessorsConfigInjection(CustomTestCase):
"""Regression tests for processors that override process_mm_data.""" """Regression tests for processors that override process_mm_data."""
def _make_override_processor(self, processor_cls, mm_process_config): def _make_override_processor(self, processor_cls, mm_process_config):
@@ -461,7 +521,6 @@ class TestOverrideProcessorsConfigInjection(unittest.TestCase):
server_args.mm_process_config = mm_process_config server_args.mm_process_config = mm_process_config
server_args.mm_feature_transport = "cpu" server_args.mm_feature_transport = "cpu"
server_args.disable_fast_image_processor = True server_args.disable_fast_image_processor = True
server_args.keep_mm_feature_on_device = False
server_args.skip_tokenizer_init = False server_args.skip_tokenizer_init = False
mock_hf_processor = MagicMock() mock_hf_processor = MagicMock()
@@ -474,7 +533,6 @@ class TestOverrideProcessorsConfigInjection(unittest.TestCase):
proc = processor_cls() proc = processor_cls()
proc.server_args = server_args proc.server_args = server_args
proc.keep_mm_feature_on_device = server_args.keep_mm_feature_on_device
proc.mm_feature_transport = server_args.mm_feature_transport proc.mm_feature_transport = server_args.mm_feature_transport
proc.use_cuda_ipc = False proc.use_cuda_ipc = False
proc.disable_fast_image_processor = server_args.disable_fast_image_processor proc.disable_fast_image_processor = server_args.disable_fast_image_processor
@@ -543,7 +601,7 @@ class TestOverrideProcessorsConfigInjection(unittest.TestCase):
self.assertTrue(audio_kw.get("truncation")) self.assertTrue(audio_kw.get("truncation"))
class TestQwenVideoConfigRouting(unittest.TestCase): class TestQwenVideoConfigRouting(CustomTestCase):
def test_preprocessed_video_drops_sglang_owned_config(self): def test_preprocessed_video_drops_sglang_owned_config(self):
from sglang.srt.multimodal.processors.qwen_vl import ( from sglang.srt.multimodal.processors.qwen_vl import (
_get_processor_video_config, _get_processor_video_config,
@@ -572,7 +630,7 @@ class TestQwenVideoConfigRouting(unittest.TestCase):
self.assertIsNone(_get_processor_video_config(video_config, [None])) self.assertIsNone(_get_processor_video_config(video_config, [None]))
class TestDoubleBosGuard(unittest.TestCase): class TestDoubleBosGuard(CustomTestCase):
"""Regression test for the multimodal double-BOS bug. """Regression test for the multimodal double-BOS bug.
Repro condition (Cohere2 / Llama3-LLaVA-Next family): Repro condition (Cohere2 / Llama3-LLaVA-Next family):
@@ -595,7 +653,6 @@ class TestDoubleBosGuard(unittest.TestCase):
server_args.mm_io_worker_num = 0 server_args.mm_io_worker_num = 0
server_args.mm_feature_transport = "cpu" server_args.mm_feature_transport = "cpu"
server_args.disable_fast_image_processor = True server_args.disable_fast_image_processor = True
server_args.keep_mm_feature_on_device = True
mock_hf_processor = MagicMock() mock_hf_processor = MagicMock()
mock_hf_processor.__class__.__name__ = "TestProcessor" mock_hf_processor.__class__.__name__ = "TestProcessor"
@@ -13,10 +13,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 asyncio
import concurrent.futures
import io import io
import unittest import unittest
from unittest.mock import Mock, patch
import numpy as np import numpy as np
import requests
from PIL import Image from PIL import Image
from sglang.srt.managers.schedule_batch import Modality from sglang.srt.managers.schedule_batch import Modality
@@ -30,6 +34,9 @@ class _StubProcessor(BaseMultimodalProcessor):
# are never called: we only invoke the _load_single_item classmethod. # are never called: we only invoke the _load_single_item classmethod.
gpu_image_decode = False gpu_image_decode = False
async def process_mm_data_async(self, *args, **kwargs):
raise NotImplementedError
def _png_bytes(mode: str = "RGB", size=(8, 8)) -> bytes: def _png_bytes(mode: str = "RGB", size=(8, 8)) -> bytes:
arr = (np.random.RandomState(0).rand(size[1], size[0], 3) * 255).astype("uint8") arr = (np.random.RandomState(0).rand(size[1], size[0], 3) * 255).astype("uint8")
@@ -75,6 +82,45 @@ class TestLoadSingleItemImageDecode(CustomTestCase):
ref = Image.open(io.BytesIO(data)).convert("RGB") ref = Image.open(io.BytesIO(data)).convert("RGB")
np.testing.assert_array_equal(np.asarray(img), np.asarray(ref)) np.testing.assert_array_equal(np.asarray(img), np.asarray(ref))
def test_fast_loader_preserves_invalid_input_as_value_error(self):
processor = object.__new__(_StubProcessor)
future = concurrent.futures.Future()
future.set_exception(ValueError("invalid base64 image"))
processor._submit_mm_data_loading_tasks_simple = Mock(
side_effect=[[(Modality.IMAGE, 0, future)], [], []]
)
with self.assertRaisesRegex(ValueError, "invalid base64 image"):
asyncio.run(
processor.fast_load_mm_data(
prompt="<image>",
multimodal_tokens=Mock(),
image_data=["bad-image"],
)
)
def test_unreachable_image_url_is_a_client_error(self):
with patch(
"sglang.srt.multimodal.processors.base_processor.load_image",
side_effect=requests.ConnectionError("connection refused"),
):
with self.assertRaisesRegex(ValueError, "connection refused"):
_StubProcessor._load_single_item(
"https://127.0.0.1:1/not-an-image.png", Modality.IMAGE
)
def test_invalid_image_bytes_are_a_client_error(self):
with self.assertRaisesRegex(ValueError, "cannot identify image file"):
_StubProcessor._load_single_item(b"not an image", Modality.IMAGE)
def test_unexpected_loader_bug_remains_a_server_error(self):
with patch(
"sglang.srt.multimodal.processors.base_processor.load_image",
side_effect=TypeError("unexpected loader bug"),
):
with self.assertRaisesRegex(RuntimeError, "unexpected loader bug"):
_StubProcessor._load_single_item(b"image", Modality.IMAGE)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()