vlm: parallelize multimodal preprocessing with customized worker num (#31438)
This commit is contained in:
@@ -19,6 +19,7 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
MultimodalInputFormat,
|
MultimodalInputFormat,
|
||||||
MultimodalProcessorOutput,
|
MultimodalProcessorOutput,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecutor
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
envs,
|
envs,
|
||||||
@@ -180,6 +181,9 @@ 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
|
||||||
|
auto_mm_processor_worker_num = 1
|
||||||
|
auto_mm_io_worker_num = 4
|
||||||
|
supports_mm_processor_concurrency = False
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, hf_config, server_args, _processor, transport_mode, *args, **kwargs
|
self, hf_config, server_args, _processor, transport_mode, *args, **kwargs
|
||||||
@@ -222,8 +226,63 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
# FIXME: not accurate, model and image specific
|
# FIXME: not accurate, model and image specific
|
||||||
self.NUM_TOKEN_PER_FRAME = 330
|
self.NUM_TOKEN_PER_FRAME = 330
|
||||||
|
|
||||||
|
requested_mm_io_worker_num = self.server_args.mm_io_worker_num
|
||||||
|
env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS")
|
||||||
|
if requested_mm_io_worker_num:
|
||||||
|
self.mm_io_worker_num = requested_mm_io_worker_num
|
||||||
|
io_worker_source = "explicit"
|
||||||
|
elif env_mm_io_worker_num is not None:
|
||||||
|
self.mm_io_worker_num = int(env_mm_io_worker_num)
|
||||||
|
io_worker_source = "environment"
|
||||||
|
else:
|
||||||
|
self.mm_io_worker_num = self.auto_mm_io_worker_num
|
||||||
|
io_worker_source = "auto"
|
||||||
self.io_executor = concurrent.futures.ThreadPoolExecutor(
|
self.io_executor = concurrent.futures.ThreadPoolExecutor(
|
||||||
max_workers=int(os.environ.get("SGLANG_IO_WORKERS", 4))
|
max_workers=self.mm_io_worker_num,
|
||||||
|
thread_name_prefix="sglang-mm-io",
|
||||||
|
)
|
||||||
|
if self.mm_io_worker_num > 4:
|
||||||
|
logger.info(
|
||||||
|
"Multimodal data loading enabled with %d worker threads (%s).",
|
||||||
|
self.mm_io_worker_num,
|
||||||
|
io_worker_source,
|
||||||
|
)
|
||||||
|
skip_mm_pool = kwargs.get("skip_mm_pool", False)
|
||||||
|
requested_mm_processor_worker_num = self.server_args.mm_processor_worker_num
|
||||||
|
self.mm_processor_worker_num = (
|
||||||
|
1
|
||||||
|
if skip_mm_pool
|
||||||
|
else requested_mm_processor_worker_num or self.auto_mm_processor_worker_num
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
self.mm_processor_worker_num > 1
|
||||||
|
and not self.supports_mm_processor_concurrency
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Concurrent multimodal processing is not supported by %s; "
|
||||||
|
"using synchronous processing.",
|
||||||
|
type(self).__name__,
|
||||||
|
)
|
||||||
|
self.mm_processor_worker_num = 1
|
||||||
|
self.mm_processor_executor = None
|
||||||
|
if self.mm_processor_worker_num > 1:
|
||||||
|
try:
|
||||||
|
self.mm_processor_executor = MultimodalProcessorExecutor(
|
||||||
|
self._processor, self.mm_processor_worker_num
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Unable to clone the multimodal processor for concurrent "
|
||||||
|
"workers; falling back to synchronous processing.",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
self.mm_processor_worker_num = 1
|
||||||
|
if self.mm_processor_executor is not None:
|
||||||
|
logger.info(
|
||||||
|
"Multimodal processor concurrency enabled with %d isolated "
|
||||||
|
"worker threads (%s).",
|
||||||
|
self.mm_processor_worker_num,
|
||||||
|
"auto" if requested_mm_processor_worker_num == 0 else "explicit",
|
||||||
)
|
)
|
||||||
self.cpu_executor = concurrent.futures.ProcessPoolExecutor(
|
self.cpu_executor = concurrent.futures.ProcessPoolExecutor(
|
||||||
mp_context=mp.get_context("fork"),
|
mp_context=mp.get_context("fork"),
|
||||||
@@ -274,8 +333,6 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
"input_features",
|
"input_features",
|
||||||
]
|
]
|
||||||
|
|
||||||
skip_mm_pool = kwargs.get("skip_mm_pool", False)
|
|
||||||
|
|
||||||
if self.use_cuda_ipc and not skip_mm_pool:
|
if self.use_cuda_ipc and not skip_mm_pool:
|
||||||
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
|
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
|
||||||
# tokenizer workers. Each worker gets an equal share so that adding
|
# tokenizer workers. Each worker gets an equal share so that adding
|
||||||
@@ -420,12 +477,25 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
video_token_id=getattr(self, "VIDEO_TOKEN_ID", None),
|
video_token_id=getattr(self, "VIDEO_TOKEN_ID", None),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _resolve_processor(self, processor=None):
|
||||||
|
if processor is None:
|
||||||
|
return self._processor, self._tokenizer
|
||||||
|
return processor, processor.tokenizer
|
||||||
|
|
||||||
def process_mm_data(
|
def process_mm_data(
|
||||||
self, input_text, images=None, videos=None, audios=None, **kwargs
|
self,
|
||||||
|
input_text,
|
||||||
|
images=None,
|
||||||
|
videos=None,
|
||||||
|
audios=None,
|
||||||
|
processor=None,
|
||||||
|
**kwargs,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""
|
"""
|
||||||
process multimodal data with transformers AutoProcessor
|
process multimodal data with transformers AutoProcessor
|
||||||
"""
|
"""
|
||||||
|
processor, tokenizer = self._resolve_processor(processor)
|
||||||
|
|
||||||
if images:
|
if images:
|
||||||
kwargs["images"] = images
|
kwargs["images"] = images
|
||||||
if self.image_config:
|
if self.image_config:
|
||||||
@@ -435,7 +505,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
if self.video_config:
|
if self.video_config:
|
||||||
kwargs.setdefault("videos_kwargs", {}).update(self.video_config)
|
kwargs.setdefault("videos_kwargs", {}).update(self.video_config)
|
||||||
if audios:
|
if audios:
|
||||||
if self._processor.__class__.__name__ in {
|
if processor.__class__.__name__ in {
|
||||||
"Gemma3nProcessor",
|
"Gemma3nProcessor",
|
||||||
"Gemma4Processor",
|
"Gemma4Processor",
|
||||||
"Gemma4UnifiedProcessor",
|
"Gemma4UnifiedProcessor",
|
||||||
@@ -453,7 +523,6 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
if self.audio_config:
|
if self.audio_config:
|
||||||
kwargs.setdefault("audio_kwargs", {}).update(self.audio_config)
|
kwargs.setdefault("audio_kwargs", {}).update(self.audio_config)
|
||||||
|
|
||||||
processor = self._processor
|
|
||||||
if (
|
if (
|
||||||
hasattr(processor, "image_processor")
|
hasattr(processor, "image_processor")
|
||||||
and isinstance(processor.image_processor, BaseImageProcessor)
|
and isinstance(processor.image_processor, BaseImageProcessor)
|
||||||
@@ -487,7 +556,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
|
|
||||||
# Avoid double BOS when the chat template already wrote one.
|
# Avoid double BOS when the chat template already wrote one.
|
||||||
if self._tokenizer_auto_adds_specials and isinstance(input_text, str):
|
if self._tokenizer_auto_adds_specials and isinstance(input_text, str):
|
||||||
bos = getattr(self._tokenizer, "bos_token", None)
|
bos = getattr(tokenizer, "bos_token", None)
|
||||||
if bos and input_text.startswith(bos):
|
if bos and input_text.startswith(bos):
|
||||||
kwargs.setdefault("add_special_tokens", False)
|
kwargs.setdefault("add_special_tokens", False)
|
||||||
|
|
||||||
@@ -1213,7 +1282,13 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
return list(items.values())
|
return list(items.values())
|
||||||
|
|
||||||
def _process_and_collect_mm_items(
|
def _process_and_collect_mm_items(
|
||||||
self, input_text: str, images=None, audios=None, videos=None, **kwargs
|
self,
|
||||||
|
input_text: str,
|
||||||
|
images=None,
|
||||||
|
audios=None,
|
||||||
|
videos=None,
|
||||||
|
processor=None,
|
||||||
|
**kwargs,
|
||||||
) -> Tuple[List[MultimodalDataItem], torch.Tensor, dict]:
|
) -> Tuple[List[MultimodalDataItem], torch.Tensor, dict]:
|
||||||
"""
|
"""
|
||||||
Helper method to process multimodal data and create mm_items in one step.
|
Helper method to process multimodal data and create mm_items in one step.
|
||||||
@@ -1221,8 +1296,14 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of (created mm_items, input_ids)
|
Tuple of (created mm_items, input_ids)
|
||||||
"""
|
"""
|
||||||
|
if processor is not None:
|
||||||
|
kwargs["processor"] = processor
|
||||||
ret = self.process_mm_data(
|
ret = self.process_mm_data(
|
||||||
input_text=input_text, images=images, audios=audios, videos=videos, **kwargs
|
input_text=input_text,
|
||||||
|
images=images,
|
||||||
|
audios=audios,
|
||||||
|
videos=videos,
|
||||||
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
input_ids = ret["input_ids"].flatten()
|
input_ids = ret["input_ids"].flatten()
|
||||||
@@ -1321,6 +1402,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self,
|
self,
|
||||||
base_output: BaseMultiModalProcessorOutput,
|
base_output: BaseMultiModalProcessorOutput,
|
||||||
mm_tokens: MultimodalSpecialTokens,
|
mm_tokens: MultimodalSpecialTokens,
|
||||||
|
processor=None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tuple[List[MultimodalDataItem], torch.Tensor, dict]:
|
) -> Tuple[List[MultimodalDataItem], torch.Tensor, dict]:
|
||||||
"""
|
"""
|
||||||
@@ -1330,11 +1412,14 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of (list of mm_items, input_ids)
|
Tuple of (list of mm_items, input_ids)
|
||||||
"""
|
"""
|
||||||
|
processor_override = processor
|
||||||
|
processor, tokenizer = self._resolve_processor(processor)
|
||||||
|
|
||||||
# Collect all items and categorize them
|
# Collect all items and categorize them
|
||||||
all_loaded_data = base_output.organize_results()
|
all_loaded_data = base_output.organize_results()
|
||||||
# Handle text-only case
|
# Handle text-only case
|
||||||
if not all_loaded_data:
|
if not all_loaded_data:
|
||||||
input_ids = self._tokenizer(
|
input_ids = tokenizer(
|
||||||
base_output.input_text,
|
base_output.input_text,
|
||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
add_special_tokens=True,
|
add_special_tokens=True,
|
||||||
@@ -1358,6 +1443,8 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
input_ids = None
|
input_ids = None
|
||||||
# Handle raw items (need processing)
|
# Handle raw items (need processing)
|
||||||
if raw_images or raw_audios or raw_videos:
|
if raw_images or raw_audios or raw_videos:
|
||||||
|
if processor_override is not None:
|
||||||
|
kwargs["processor"] = processor
|
||||||
collected_items, input_ids, ret = self._process_and_collect_mm_items(
|
collected_items, input_ids, ret = self._process_and_collect_mm_items(
|
||||||
input_text=base_output.input_text,
|
input_text=base_output.input_text,
|
||||||
images=raw_images,
|
images=raw_images,
|
||||||
@@ -1447,7 +1534,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
break
|
break
|
||||||
|
|
||||||
if input_ids is None:
|
if input_ids is None:
|
||||||
input_ids = self._tokenizer(
|
input_ids = tokenizer(
|
||||||
base_output.input_text,
|
base_output.input_text,
|
||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
add_special_tokens=True,
|
add_special_tokens=True,
|
||||||
@@ -1496,3 +1583,20 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
return all_collected_items, input_ids, ret
|
return all_collected_items, input_ids, ret
|
||||||
|
|
||||||
|
async def process_and_combine_mm_data_async(
|
||||||
|
self,
|
||||||
|
base_output: BaseMultiModalProcessorOutput,
|
||||||
|
mm_tokens: MultimodalSpecialTokens,
|
||||||
|
**kwargs,
|
||||||
|
) -> Tuple[List[MultimodalDataItem], torch.Tensor, dict]:
|
||||||
|
"""Run multimodal preprocessing without blocking the event loop."""
|
||||||
|
if self.mm_processor_executor is None:
|
||||||
|
return self.process_and_combine_mm_data(base_output, mm_tokens, **kwargs)
|
||||||
|
|
||||||
|
return await self.mm_processor_executor.run(
|
||||||
|
self.process_and_combine_mm_data,
|
||||||
|
base_output,
|
||||||
|
mm_tokens,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
|
import copy
|
||||||
|
import threading
|
||||||
|
from typing import Any, Callable, TypeVar
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class _WorkerState(threading.local):
|
||||||
|
def __init__(self):
|
||||||
|
self.processor = None
|
||||||
|
|
||||||
|
|
||||||
|
class MultimodalProcessorExecutor:
|
||||||
|
"""Run processor calls on isolated, thread-local processor clones."""
|
||||||
|
|
||||||
|
def __init__(self, processor: Any, max_workers: int):
|
||||||
|
self._processor = processor
|
||||||
|
self._processor_clones = [copy.deepcopy(processor) for _ in range(max_workers)]
|
||||||
|
self._executor = concurrent.futures.ThreadPoolExecutor(
|
||||||
|
max_workers=max_workers,
|
||||||
|
thread_name_prefix="sglang-mm-processor",
|
||||||
|
)
|
||||||
|
self._worker_state = _WorkerState()
|
||||||
|
self._clone_lock = threading.Lock()
|
||||||
|
|
||||||
|
async def run(self, function: Callable[..., T], *args: Any, **kwargs: Any) -> T:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
return await loop.run_in_executor(
|
||||||
|
self._executor, self._run, function, args, kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def _run(
|
||||||
|
self,
|
||||||
|
function: Callable[..., T],
|
||||||
|
args: tuple[Any, ...],
|
||||||
|
kwargs: dict[str, Any],
|
||||||
|
) -> T:
|
||||||
|
processor = self._worker_state.processor
|
||||||
|
if processor is None:
|
||||||
|
with self._clone_lock:
|
||||||
|
processor = (
|
||||||
|
self._processor_clones.pop()
|
||||||
|
if self._processor_clones
|
||||||
|
else copy.deepcopy(self._processor)
|
||||||
|
)
|
||||||
|
self._worker_state.processor = processor
|
||||||
|
return function(*args, processor=processor, **kwargs)
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
self._executor.shutdown()
|
||||||
@@ -273,6 +273,22 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
|
|
||||||
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
|
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
|
||||||
self.model_type = hf_config.model_type
|
self.model_type = hf_config.model_type
|
||||||
|
if self.model_type in (
|
||||||
|
"qwen2_vl",
|
||||||
|
"qwen2_5_vl",
|
||||||
|
"qwen3_vl",
|
||||||
|
"qwen3_vl_moe",
|
||||||
|
"qwen3_5",
|
||||||
|
"qwen3_5_moe",
|
||||||
|
"intern_s2_preview",
|
||||||
|
):
|
||||||
|
# Two workers overlap CPU preprocessing without over-fragmenting
|
||||||
|
# burst arrivals into smaller GPU prefill batches. Higher counts can
|
||||||
|
# improve short-output TTFT, but regress long-output throughput on
|
||||||
|
# Blackwell when requests reach the scheduler too far apart.
|
||||||
|
self.auto_mm_processor_worker_num = 2
|
||||||
|
self.auto_mm_io_worker_num = 16
|
||||||
|
self.supports_mm_processor_concurrency = True
|
||||||
if hf_config.model_type == "qwen3_omni_moe":
|
if hf_config.model_type == "qwen3_omni_moe":
|
||||||
hf_config = hf_config.thinker_config
|
hf_config = hf_config.thinker_config
|
||||||
|
|
||||||
@@ -711,14 +727,14 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
"qwen3_5_moe",
|
"qwen3_5_moe",
|
||||||
"intern_s2_preview",
|
"intern_s2_preview",
|
||||||
):
|
):
|
||||||
mm_items, input_ids, ret = self.process_and_combine_mm_data(
|
mm_items, input_ids, ret = await self.process_and_combine_mm_data_async(
|
||||||
base_output,
|
base_output,
|
||||||
self.mm_tokens,
|
self.mm_tokens,
|
||||||
video_metadata=video_metadata,
|
video_metadata=video_metadata,
|
||||||
do_sample_frames=False,
|
do_sample_frames=False,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
mm_items, input_ids, ret = self.process_and_combine_mm_data(
|
mm_items, input_ids, ret = await self.process_and_combine_mm_data_async(
|
||||||
base_output, self.mm_tokens
|
base_output, self.mm_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -2283,6 +2283,18 @@ class ServerArgs:
|
|||||||
type_parser=json.loads,
|
type_parser=json.loads,
|
||||||
),
|
),
|
||||||
] = None
|
] = None
|
||||||
|
mm_processor_worker_num: A[
|
||||||
|
int,
|
||||||
|
"Number of threads for multimodal processor calls. 0 selects the "
|
||||||
|
"model-specific default. Only processors with isolated-worker support "
|
||||||
|
"can use more than one thread.",
|
||||||
|
] = 0
|
||||||
|
mm_io_worker_num: A[
|
||||||
|
int,
|
||||||
|
"Number of threads for multimodal data loading and decoding. 0 selects "
|
||||||
|
"the model-specific default. SGLANG_IO_WORKERS remains supported as an "
|
||||||
|
"environment override when this argument is 0.",
|
||||||
|
] = 0
|
||||||
limit_mm_data_per_request: A[
|
limit_mm_data_per_request: A[
|
||||||
Optional[Union[str, Dict[str, int]]],
|
Optional[Union[str, Dict[str, int]]],
|
||||||
Arg(
|
Arg(
|
||||||
@@ -7630,6 +7642,10 @@ class ServerArgs:
|
|||||||
|
|
||||||
assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1"
|
assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1"
|
||||||
assert self.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1"
|
assert self.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1"
|
||||||
|
assert (
|
||||||
|
self.mm_processor_worker_num >= 0
|
||||||
|
), "Multimodal processor worker num must >= 0"
|
||||||
|
assert self.mm_io_worker_num >= 0, "Multimodal I/O worker num must >= 0"
|
||||||
self.validate_buckets_rule(
|
self.validate_buckets_rule(
|
||||||
"--prompt-tokens-buckets", self.prompt_tokens_buckets
|
"--prompt-tokens-buckets", self.prompt_tokens_buckets
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
import os
|
||||||
|
import threading
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
@@ -64,7 +66,12 @@ class TestMmProcessConfigValidation(unittest.TestCase):
|
|||||||
class TestBaseProcessorConfigExtraction(unittest.TestCase):
|
class TestBaseProcessorConfigExtraction(unittest.TestCase):
|
||||||
"""Verify BaseMultimodalProcessor.__init__ extracts configs from server_args."""
|
"""Verify BaseMultimodalProcessor.__init__ extracts configs from server_args."""
|
||||||
|
|
||||||
def _make_processor(self, mm_process_config):
|
def _make_processor(
|
||||||
|
self,
|
||||||
|
mm_process_config,
|
||||||
|
mm_processor_worker_num=0,
|
||||||
|
mm_io_worker_num=0,
|
||||||
|
):
|
||||||
"""Create a BaseMultimodalProcessor via the real __init__ with mocked deps."""
|
"""Create a BaseMultimodalProcessor via the real __init__ with mocked deps."""
|
||||||
from sglang.srt.multimodal.processors.base_processor import (
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
BaseMultimodalProcessor,
|
BaseMultimodalProcessor,
|
||||||
@@ -72,6 +79,8 @@ class TestBaseProcessorConfigExtraction(unittest.TestCase):
|
|||||||
|
|
||||||
server_args = MagicMock()
|
server_args = MagicMock()
|
||||||
server_args.mm_process_config = mm_process_config
|
server_args.mm_process_config = mm_process_config
|
||||||
|
server_args.mm_processor_worker_num = mm_processor_worker_num
|
||||||
|
server_args.mm_io_worker_num = mm_io_worker_num
|
||||||
|
|
||||||
hf_config = MagicMock()
|
hf_config = MagicMock()
|
||||||
mock_hf_processor = MagicMock()
|
mock_hf_processor = MagicMock()
|
||||||
@@ -103,6 +112,52 @@ class TestBaseProcessorConfigExtraction(unittest.TestCase):
|
|||||||
self.assertEqual(proc.video_config, {})
|
self.assertEqual(proc.video_config, {})
|
||||||
self.assertEqual(proc.audio_config, {})
|
self.assertEqual(proc.audio_config, {})
|
||||||
|
|
||||||
|
def test_model_specific_auto_worker_count_enables_executor(self):
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.dict(os.environ, {}, clear=False):
|
||||||
|
os.environ.pop("SGLANG_IO_WORKERS", None)
|
||||||
|
with patch.object(
|
||||||
|
BaseMultimodalProcessor, "auto_mm_processor_worker_num", 4
|
||||||
|
), patch.object(
|
||||||
|
BaseMultimodalProcessor, "auto_mm_io_worker_num", 16
|
||||||
|
), patch.object(
|
||||||
|
BaseMultimodalProcessor, "supports_mm_processor_concurrency", True
|
||||||
|
):
|
||||||
|
proc = self._make_processor({})
|
||||||
|
try:
|
||||||
|
self.assertEqual(proc.mm_processor_worker_num, 4)
|
||||||
|
self.assertEqual(proc.mm_io_worker_num, 16)
|
||||||
|
self.assertIsNotNone(proc.mm_processor_executor)
|
||||||
|
finally:
|
||||||
|
proc.mm_processor_executor.shutdown()
|
||||||
|
|
||||||
|
def test_explicit_single_worker_disables_executor(self):
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(BaseMultimodalProcessor, "auto_mm_processor_worker_num", 4):
|
||||||
|
proc = self._make_processor({}, mm_processor_worker_num=1)
|
||||||
|
self.assertEqual(proc.mm_processor_worker_num, 1)
|
||||||
|
self.assertIsNone(proc.mm_processor_executor)
|
||||||
|
|
||||||
|
def test_parallel_workers_require_processor_support(self):
|
||||||
|
proc = self._make_processor({}, mm_processor_worker_num=2)
|
||||||
|
self.assertEqual(proc.mm_processor_worker_num, 1)
|
||||||
|
self.assertIsNone(proc.mm_processor_executor)
|
||||||
|
|
||||||
|
def test_explicit_io_worker_count_overrides_auto(self):
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(BaseMultimodalProcessor, "auto_mm_io_worker_num", 16):
|
||||||
|
proc = self._make_processor({}, mm_io_worker_num=6)
|
||||||
|
self.assertEqual(proc.mm_io_worker_num, 6)
|
||||||
|
|
||||||
|
|
||||||
class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
|
class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -113,6 +168,8 @@ class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
|
|||||||
disable_fast_image_processor=False,
|
disable_fast_image_processor=False,
|
||||||
skip_tokenizer_init=False,
|
skip_tokenizer_init=False,
|
||||||
mm_process_config={},
|
mm_process_config={},
|
||||||
|
mm_processor_worker_num=0,
|
||||||
|
mm_io_worker_num=0,
|
||||||
tokenizer_worker_num=1,
|
tokenizer_worker_num=1,
|
||||||
base_gpu_id=2,
|
base_gpu_id=2,
|
||||||
)
|
)
|
||||||
@@ -160,6 +217,100 @@ class TestMultimodalFeatureTransportRuntime(unittest.TestCase):
|
|||||||
memory_pool.assert_not_called()
|
memory_pool.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_dedicated_executor_runs_processor_off_event_loop(self):
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
from sglang.srt.multimodal.processors.executor import (
|
||||||
|
MultimodalProcessorExecutor,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
BaseMultimodalProcessor, "__abstractmethods__", set()
|
||||||
|
), patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
|
||||||
|
processor = BaseMultimodalProcessor()
|
||||||
|
|
||||||
|
processor.mm_processor_executor = MultimodalProcessorExecutor(
|
||||||
|
SimpleNamespace(tokenizer=object()), max_workers=2
|
||||||
|
)
|
||||||
|
processor.process_and_combine_mm_data = MagicMock(
|
||||||
|
side_effect=lambda *_args, **_kwargs: threading.current_thread().name
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
thread_name = await processor.process_and_combine_mm_data_async(
|
||||||
|
MagicMock(), MagicMock(), marker=True
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
processor.mm_processor_executor.shutdown()
|
||||||
|
|
||||||
|
self.assertTrue(thread_name.startswith("sglang-mm-processor"))
|
||||||
|
processor.process_and_combine_mm_data.assert_called_once()
|
||||||
|
self.assertTrue(
|
||||||
|
processor.process_and_combine_mm_data.call_args.kwargs["marker"]
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_single_worker_preserves_synchronous_path(self):
|
||||||
|
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.mm_processor_executor = None
|
||||||
|
processor.process_and_combine_mm_data = MagicMock(return_value="synchronous")
|
||||||
|
|
||||||
|
result = await processor.process_and_combine_mm_data_async(
|
||||||
|
MagicMock(), MagicMock()
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(result, "synchronous")
|
||||||
|
processor.process_and_combine_mm_data.assert_called_once()
|
||||||
|
|
||||||
|
async def test_worker_reuses_precreated_private_processor_clone(self):
|
||||||
|
from sglang.srt.multimodal.processors.executor import (
|
||||||
|
MultimodalProcessorExecutor,
|
||||||
|
)
|
||||||
|
|
||||||
|
executor = MultimodalProcessorExecutor(object(), max_workers=2)
|
||||||
|
return_processor = lambda *, processor: processor
|
||||||
|
try:
|
||||||
|
first = await executor.run(return_processor)
|
||||||
|
second = await executor.run(return_processor)
|
||||||
|
finally:
|
||||||
|
executor.shutdown()
|
||||||
|
|
||||||
|
self.assertIs(first, second)
|
||||||
|
|
||||||
|
async def test_replacement_worker_lazily_clones_processor(self):
|
||||||
|
from sglang.srt.multimodal.processors import executor as executor_module
|
||||||
|
|
||||||
|
source_processor = object()
|
||||||
|
replacement_clone = SimpleNamespace(tokenizer=object())
|
||||||
|
with patch.object(
|
||||||
|
executor_module.copy,
|
||||||
|
"deepcopy",
|
||||||
|
return_value=replacement_clone,
|
||||||
|
) as deepcopy:
|
||||||
|
executor = executor_module.MultimodalProcessorExecutor(
|
||||||
|
source_processor, max_workers=2
|
||||||
|
)
|
||||||
|
executor._processor_clones.clear()
|
||||||
|
return_processor = lambda *, processor: processor
|
||||||
|
try:
|
||||||
|
first = await executor.run(return_processor)
|
||||||
|
second = await executor.run(return_processor)
|
||||||
|
finally:
|
||||||
|
executor.shutdown()
|
||||||
|
|
||||||
|
self.assertIs(first, replacement_clone)
|
||||||
|
self.assertIs(first, second)
|
||||||
|
self.assertEqual(deepcopy.call_count, 3)
|
||||||
|
|
||||||
|
|
||||||
class TestProcessMmDataKwargs(unittest.TestCase):
|
class TestProcessMmDataKwargs(unittest.TestCase):
|
||||||
"""Verify process_mm_data injects per-modality kwargs correctly."""
|
"""Verify process_mm_data injects per-modality kwargs correctly."""
|
||||||
|
|
||||||
@@ -380,6 +531,8 @@ class TestDoubleBosGuard(unittest.TestCase):
|
|||||||
|
|
||||||
server_args = MagicMock()
|
server_args = MagicMock()
|
||||||
server_args.mm_process_config = {}
|
server_args.mm_process_config = {}
|
||||||
|
server_args.mm_processor_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
|
server_args.keep_mm_feature_on_device = True
|
||||||
|
|||||||
Reference in New Issue
Block a user