VLM: support passing --mm-process-config for all models (#18467)
This commit is contained in:
@@ -143,7 +143,7 @@ Use this flag when you have sufficient GPU memory and want to minimize latency f
|
|||||||
|
|
||||||
- **Use `--mm-process-config '{"image":{"max_pixels":1048576},"video":{"fps":3,"max_pixels":602112,"max_frames":60}}'`**: To set `image`, `video`, and `audio` input limits.
|
- **Use `--mm-process-config '{"image":{"max_pixels":1048576},"video":{"fps":3,"max_pixels":602112,"max_frames":60}}'`**: To set `image`, `video`, and `audio` input limits.
|
||||||
|
|
||||||
This can reduce GPU memory usage, improve inference speed, and help to avoid OOM, but may impact model performance, thus set a proper value based on your specific use case. Currently, only `qwen_vl` supports this config. Please refer to [qwen_vl processor](https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/multimodal/processors/qwen_vl.py) for understanding the meaning of each parameter.
|
This can reduce GPU memory usage, improve inference speed, and help to avoid OOM, but may impact model performance, thus set a proper value based on your specific use case. The config entries are passed as `images_kwargs`, `videos_kwargs`, and `audio_kwargs` to the HuggingFace processor, so each modality's settings are kept separate and do not collide. Refer to the HuggingFace documentation for your model's processor to understand the available parameters.
|
||||||
|
|
||||||
### Bidirectional Attention in Multimodal Model Serving
|
### Bidirectional Attention in Multimodal Model Serving
|
||||||
**Note for serving the Gemma-3 multimodal model**:
|
**Note for serving the Gemma-3 multimodal model**:
|
||||||
|
|||||||
@@ -184,6 +184,11 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.transport_mode = transport_mode
|
self.transport_mode = transport_mode
|
||||||
|
|
||||||
|
mm_process_config = self.server_args.mm_process_config
|
||||||
|
self.image_config = mm_process_config.get("image", {})
|
||||||
|
self.video_config = mm_process_config.get("video", {})
|
||||||
|
self.audio_config = mm_process_config.get("audio", {})
|
||||||
|
|
||||||
# Resolve tokenizer: some processors (e.g. InternVL) pass a tokenizer
|
# Resolve tokenizer: some processors (e.g. InternVL) pass a tokenizer
|
||||||
# directly as _processor rather than a processor that wraps a tokenizer.
|
# directly as _processor rather than a processor that wraps a tokenizer.
|
||||||
if hasattr(self._processor, "tokenizer"):
|
if hasattr(self._processor, "tokenizer"):
|
||||||
@@ -381,8 +386,12 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
"""
|
"""
|
||||||
if images:
|
if images:
|
||||||
kwargs["images"] = images
|
kwargs["images"] = images
|
||||||
|
if self.image_config:
|
||||||
|
kwargs.setdefault("images_kwargs", {}).update(self.image_config)
|
||||||
if videos:
|
if videos:
|
||||||
kwargs["videos"] = videos
|
kwargs["videos"] = videos
|
||||||
|
if self.video_config:
|
||||||
|
kwargs.setdefault("videos_kwargs", {}).update(self.video_config)
|
||||||
if audios:
|
if audios:
|
||||||
if self._processor.__class__.__name__ in {
|
if self._processor.__class__.__name__ in {
|
||||||
"Gemma3nProcessor",
|
"Gemma3nProcessor",
|
||||||
@@ -394,10 +403,12 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
}:
|
}:
|
||||||
# Note(Xinyuan): for gemma3n, ref: https://github.com/huggingface/transformers/blob/ccf2ca162e33f381e454cdb74bf4b41a51ab976d/src/transformers/models/gemma3n/processing_gemma3n.py#L107
|
# Note(Xinyuan): for gemma3n, ref: https://github.com/huggingface/transformers/blob/ccf2ca162e33f381e454cdb74bf4b41a51ab976d/src/transformers/models/gemma3n/processing_gemma3n.py#L107
|
||||||
kwargs["audio"] = audios
|
kwargs["audio"] = audios
|
||||||
kwargs["audio_kwargs"] = {}
|
kwargs.setdefault("audio_kwargs", {})
|
||||||
kwargs["audio_kwargs"].setdefault("truncation", False)
|
kwargs["audio_kwargs"].setdefault("truncation", False)
|
||||||
else:
|
else:
|
||||||
kwargs["audios"] = audios
|
kwargs["audios"] = audios
|
||||||
|
if self.audio_config:
|
||||||
|
kwargs.setdefault("audio_kwargs", {}).update(self.audio_config)
|
||||||
|
|
||||||
processor = self._processor
|
processor = self._processor
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -292,8 +292,12 @@ class Ernie4_5_VLImageProcessor(SGLangBaseProcessor):
|
|||||||
"""
|
"""
|
||||||
if images:
|
if images:
|
||||||
kwargs["images"] = images
|
kwargs["images"] = images
|
||||||
|
if self.image_config:
|
||||||
|
kwargs.setdefault("images_kwargs", {}).update(self.image_config)
|
||||||
if videos:
|
if videos:
|
||||||
kwargs["videos"] = videos
|
kwargs["videos"] = videos
|
||||||
|
if self.video_config:
|
||||||
|
kwargs.setdefault("videos_kwargs", {}).update(self.video_config)
|
||||||
|
|
||||||
processor = self._processor
|
processor = self._processor
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -57,8 +57,10 @@ class MiDashengLMMultimodalProcessor(BaseMultimodalProcessor):
|
|||||||
kwargs["videos"] = videos
|
kwargs["videos"] = videos
|
||||||
if audios:
|
if audios:
|
||||||
kwargs["audio"] = audios
|
kwargs["audio"] = audios
|
||||||
kwargs["audio_kwargs"] = {}
|
kwargs.setdefault("audio_kwargs", {})
|
||||||
kwargs["audio_kwargs"].setdefault("truncation", False)
|
kwargs["audio_kwargs"].setdefault("truncation", False)
|
||||||
|
if self.audio_config:
|
||||||
|
kwargs["audio_kwargs"].update(self.audio_config)
|
||||||
|
|
||||||
processor = self._processor
|
processor = self._processor
|
||||||
result = processor.__call__(
|
result = processor.__call__(
|
||||||
|
|||||||
@@ -267,9 +267,6 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
self.audio_start_token_id = getattr(hf_config, "audio_start_token_id", None)
|
self.audio_start_token_id = getattr(hf_config, "audio_start_token_id", None)
|
||||||
self.audio_token_id = getattr(hf_config, "audio_token_id", None)
|
self.audio_token_id = getattr(hf_config, "audio_token_id", None)
|
||||||
|
|
||||||
self.image_config = server_args.mm_process_config.get("image", {})
|
|
||||||
self.video_config = server_args.mm_process_config.get("video", {})
|
|
||||||
|
|
||||||
self.mm_tokens = MultimodalSpecialTokens(
|
self.mm_tokens = MultimodalSpecialTokens(
|
||||||
image_token="<|vision_start|><|image_pad|><|vision_end|>",
|
image_token="<|vision_start|><|image_pad|><|vision_end|>",
|
||||||
image_token_id=hf_config.image_token_id,
|
image_token_id=hf_config.image_token_id,
|
||||||
|
|||||||
@@ -762,6 +762,8 @@ class ServerArgs:
|
|||||||
# Normalize load balancing defaults early (before dummy-model short-circuit).
|
# Normalize load balancing defaults early (before dummy-model short-circuit).
|
||||||
self._handle_load_balance_method()
|
self._handle_load_balance_method()
|
||||||
|
|
||||||
|
# Validate mm_process_config before dummy-model early return.
|
||||||
|
self._handle_multimodal()
|
||||||
# Validate SSL arguments early (before dummy-model short-circuit).
|
# Validate SSL arguments early (before dummy-model short-circuit).
|
||||||
self._handle_ssl_validation()
|
self._handle_ssl_validation()
|
||||||
|
|
||||||
@@ -938,18 +940,37 @@ class ServerArgs:
|
|||||||
"--enable-http2 requires the 'granian' package. "
|
"--enable-http2 requires the 'granian' package. "
|
||||||
'Install it with: pip install "sglang[http2]"'
|
'Install it with: pip install "sglang[http2]"'
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.enable_ssl_refresh:
|
if self.enable_ssl_refresh:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-ssl-refresh is not supported with --enable-http2. "
|
"--enable-ssl-refresh is not supported with --enable-http2. "
|
||||||
"Granian does not support SSL certificate hot-reloading. "
|
"Granian does not support SSL certificate hot-reloading. "
|
||||||
"Use Uvicorn (the default) or handle certificate rotation externally."
|
"Use Uvicorn (the default) or handle certificate rotation externally."
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.tokenizer_worker_num > 1:
|
if self.tokenizer_worker_num > 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-http2 does not yet support --tokenizer-worker-num > 1. "
|
"--enable-http2 does not yet support --tokenizer-worker-num > 1. "
|
||||||
"Multi-worker HTTP/2 support will be added in a future release."
|
"Multi-worker HTTP/2 support will be added in a future release."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _handle_multimodal(self):
|
||||||
|
"""Validate mm_process_config structure before model loading."""
|
||||||
|
if self.mm_process_config is not None:
|
||||||
|
if not isinstance(self.mm_process_config, dict):
|
||||||
|
raise TypeError(
|
||||||
|
f"mm_process_config must be a dict, "
|
||||||
|
f"but got {type(self.mm_process_config)}"
|
||||||
|
)
|
||||||
|
for key in ("image", "video", "audio"):
|
||||||
|
if key in self.mm_process_config and not isinstance(
|
||||||
|
self.mm_process_config[key], dict
|
||||||
|
):
|
||||||
|
raise TypeError(
|
||||||
|
f"mm_process_config['{key}'] must be a dict, "
|
||||||
|
f"but got {type(self.mm_process_config[key])}"
|
||||||
|
)
|
||||||
|
|
||||||
def _handle_deprecated_args(self):
|
def _handle_deprecated_args(self):
|
||||||
# Handle deprecated tool call parsers
|
# Handle deprecated tool call parsers
|
||||||
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
|
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
|
||||||
|
|||||||
@@ -0,0 +1,290 @@
|
|||||||
|
import unittest
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=5, suite="stage-b-test-1-gpu-small")
|
||||||
|
register_amd_ci(est_time=1, suite="stage-b-test-1-gpu-small-amd")
|
||||||
|
|
||||||
|
|
||||||
|
class TestMmProcessConfigValidation(unittest.TestCase):
|
||||||
|
"""Server-args validation for mm_process_config."""
|
||||||
|
|
||||||
|
def test_valid_config_accepted(self):
|
||||||
|
args = ServerArgs(
|
||||||
|
model_path="dummy",
|
||||||
|
mm_process_config={"image": {"max_pixels": 5000000}},
|
||||||
|
)
|
||||||
|
self.assertEqual(args.mm_process_config, {"image": {"max_pixels": 5000000}})
|
||||||
|
|
||||||
|
def test_empty_config_accepted(self):
|
||||||
|
args = ServerArgs(model_path="dummy", mm_process_config={})
|
||||||
|
self.assertEqual(args.mm_process_config, {})
|
||||||
|
|
||||||
|
def test_none_config_defaults_to_empty_dict(self):
|
||||||
|
args = ServerArgs(model_path="dummy", mm_process_config=None)
|
||||||
|
# None is kept as-is for dummy models (default happens after early return)
|
||||||
|
# but for real models it would be set to {}
|
||||||
|
self.assertIsNone(args.mm_process_config)
|
||||||
|
|
||||||
|
def test_top_level_non_dict_rejected(self):
|
||||||
|
with self.assertRaises(TypeError) as ctx:
|
||||||
|
ServerArgs(model_path="dummy", mm_process_config="bad")
|
||||||
|
self.assertIn("mm_process_config must be a dict", str(ctx.exception))
|
||||||
|
|
||||||
|
def test_modality_non_dict_rejected_image(self):
|
||||||
|
with self.assertRaises(TypeError) as ctx:
|
||||||
|
ServerArgs(model_path="dummy", mm_process_config={"image": "bad"})
|
||||||
|
self.assertIn("mm_process_config['image'] must be a dict", str(ctx.exception))
|
||||||
|
|
||||||
|
def test_modality_non_dict_rejected_video(self):
|
||||||
|
with self.assertRaises(TypeError) as ctx:
|
||||||
|
ServerArgs(model_path="dummy", mm_process_config={"video": 123})
|
||||||
|
self.assertIn("mm_process_config['video'] must be a dict", str(ctx.exception))
|
||||||
|
|
||||||
|
def test_modality_non_dict_rejected_audio(self):
|
||||||
|
with self.assertRaises(TypeError) as ctx:
|
||||||
|
ServerArgs(model_path="dummy", mm_process_config={"audio": [1, 2]})
|
||||||
|
self.assertIn("mm_process_config['audio'] must be a dict", str(ctx.exception))
|
||||||
|
|
||||||
|
def test_multi_modality_config_accepted(self):
|
||||||
|
config = {
|
||||||
|
"image": {"max_pixels": 1048576},
|
||||||
|
"video": {"max_pixels": 602112},
|
||||||
|
"audio": {"sample_rate": 16000},
|
||||||
|
}
|
||||||
|
args = ServerArgs(model_path="dummy", mm_process_config=config)
|
||||||
|
self.assertEqual(args.mm_process_config, config)
|
||||||
|
|
||||||
|
|
||||||
|
class TestBaseProcessorConfigExtraction(unittest.TestCase):
|
||||||
|
"""Verify BaseMultimodalProcessor.__init__ extracts configs from server_args."""
|
||||||
|
|
||||||
|
def _make_processor(self, mm_process_config):
|
||||||
|
"""Create a BaseMultimodalProcessor via the real __init__ with mocked deps."""
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
server_args = MagicMock()
|
||||||
|
server_args.mm_process_config = mm_process_config
|
||||||
|
|
||||||
|
hf_config = MagicMock()
|
||||||
|
mock_hf_processor = MagicMock()
|
||||||
|
|
||||||
|
# Call real __init__ so we test actual config extraction
|
||||||
|
with patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()):
|
||||||
|
proc = BaseMultimodalProcessor(
|
||||||
|
hf_config=hf_config,
|
||||||
|
server_args=server_args,
|
||||||
|
_processor=mock_hf_processor,
|
||||||
|
transport_mode=None,
|
||||||
|
)
|
||||||
|
return proc
|
||||||
|
|
||||||
|
def test_configs_extracted(self):
|
||||||
|
config = {
|
||||||
|
"image": {"max_pixels": 5000000},
|
||||||
|
"video": {"fps": 3},
|
||||||
|
"audio": {"sample_rate": 16000},
|
||||||
|
}
|
||||||
|
proc = self._make_processor(config)
|
||||||
|
self.assertEqual(proc.image_config, {"max_pixels": 5000000})
|
||||||
|
self.assertEqual(proc.video_config, {"fps": 3})
|
||||||
|
self.assertEqual(proc.audio_config, {"sample_rate": 16000})
|
||||||
|
|
||||||
|
def test_empty_config_yields_empty_dicts(self):
|
||||||
|
proc = self._make_processor({})
|
||||||
|
self.assertEqual(proc.image_config, {})
|
||||||
|
self.assertEqual(proc.video_config, {})
|
||||||
|
self.assertEqual(proc.audio_config, {})
|
||||||
|
|
||||||
|
|
||||||
|
class TestProcessMmDataKwargs(unittest.TestCase):
|
||||||
|
"""Verify process_mm_data injects per-modality kwargs correctly."""
|
||||||
|
|
||||||
|
def _make_base_processor(self, mm_process_config):
|
||||||
|
"""Create a BaseMultimodalProcessor with process_mm_data testable."""
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
server_args = MagicMock()
|
||||||
|
server_args.mm_process_config = mm_process_config
|
||||||
|
server_args.disable_fast_image_processor = True
|
||||||
|
server_args.keep_mm_feature_on_device = True
|
||||||
|
|
||||||
|
mock_processor = MagicMock()
|
||||||
|
mock_processor.__class__.__name__ = "TestProcessor"
|
||||||
|
# Capture kwargs passed to __call__
|
||||||
|
captured_kwargs = {}
|
||||||
|
|
||||||
|
def capture_call(**kwargs):
|
||||||
|
captured_kwargs.update(kwargs)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
mock_processor.__call__ = MagicMock(side_effect=capture_call)
|
||||||
|
|
||||||
|
with patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()):
|
||||||
|
with patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
|
||||||
|
proc = BaseMultimodalProcessor()
|
||||||
|
|
||||||
|
proc.server_args = server_args
|
||||||
|
proc._processor = mock_processor
|
||||||
|
proc.image_config = mm_process_config.get("image", {})
|
||||||
|
proc.video_config = mm_process_config.get("video", {})
|
||||||
|
proc.audio_config = mm_process_config.get("audio", {})
|
||||||
|
proc.FEATURE_NAMES = []
|
||||||
|
|
||||||
|
return proc, mock_processor, captured_kwargs
|
||||||
|
|
||||||
|
def test_images_kwargs_injected(self):
|
||||||
|
config = {"image": {"max_pixels": 5000000}}
|
||||||
|
proc, mock_proc, _ = self._make_base_processor(config)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", images=["img1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("images_kwargs"), {"max_pixels": 5000000}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_videos_kwargs_injected(self):
|
||||||
|
config = {"video": {"fps": 3, "max_frames": 60}}
|
||||||
|
proc, mock_proc, _ = self._make_base_processor(config)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", videos=["vid1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("videos_kwargs"), {"fps": 3, "max_frames": 60}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_no_collision_with_overlapping_keys(self):
|
||||||
|
"""Core test: image and video both have max_pixels but stay separate."""
|
||||||
|
config = {
|
||||||
|
"image": {"max_pixels": 1048576},
|
||||||
|
"video": {"max_pixels": 602112},
|
||||||
|
}
|
||||||
|
proc, mock_proc, _ = self._make_base_processor(config)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", images=["img1"], videos=["vid1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("images_kwargs"), {"max_pixels": 1048576}
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("videos_kwargs"), {"max_pixels": 602112}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_empty_config_no_kwargs_injected(self):
|
||||||
|
proc, mock_proc, _ = self._make_base_processor({})
|
||||||
|
|
||||||
|
proc.process_mm_data("test", images=["img1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
self.assertNotIn("images_kwargs", call_kwargs.kwargs)
|
||||||
|
|
||||||
|
def test_audio_kwargs_preserved_with_config(self):
|
||||||
|
"""audio_config merges with existing truncation=False."""
|
||||||
|
config = {"audio": {"sample_rate": 16000}}
|
||||||
|
proc, mock_proc, _ = self._make_base_processor(config)
|
||||||
|
# Simulate a processor that uses singular "audio" key
|
||||||
|
mock_proc.__class__.__name__ = "Gemma3nProcessor"
|
||||||
|
|
||||||
|
proc.process_mm_data("test", audios=["aud1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
audio_kw = call_kwargs.kwargs.get("audio_kwargs", {})
|
||||||
|
self.assertFalse(audio_kw.get("truncation", True))
|
||||||
|
self.assertEqual(audio_kw.get("sample_rate"), 16000)
|
||||||
|
|
||||||
|
|
||||||
|
class TestOverrideProcessorsConfigInjection(unittest.TestCase):
|
||||||
|
"""Regression tests for processors that override process_mm_data."""
|
||||||
|
|
||||||
|
def _make_override_processor(self, processor_cls, mm_process_config):
|
||||||
|
"""Create an override processor with mocked dependencies."""
|
||||||
|
server_args = MagicMock()
|
||||||
|
server_args.mm_process_config = mm_process_config
|
||||||
|
server_args.disable_fast_image_processor = True
|
||||||
|
server_args.keep_mm_feature_on_device = False
|
||||||
|
|
||||||
|
mock_hf_processor = MagicMock()
|
||||||
|
mock_hf_processor.__class__.__name__ = "TestProcessor"
|
||||||
|
# Ernie processor accesses result["images"] after __call__,
|
||||||
|
# so return {"images": None} to pass the None-guard safely.
|
||||||
|
mock_hf_processor.__call__ = MagicMock(return_value={"images": None})
|
||||||
|
|
||||||
|
with patch.object(processor_cls, "__init__", lambda self: None):
|
||||||
|
proc = processor_cls()
|
||||||
|
|
||||||
|
proc.server_args = server_args
|
||||||
|
proc._processor = mock_hf_processor
|
||||||
|
proc.image_config = mm_process_config.get("image", {})
|
||||||
|
proc.video_config = mm_process_config.get("video", {})
|
||||||
|
proc.audio_config = mm_process_config.get("audio", {})
|
||||||
|
proc.FEATURE_NAMES = []
|
||||||
|
|
||||||
|
return proc, mock_hf_processor
|
||||||
|
|
||||||
|
def test_ernie45_vl_injects_images_kwargs(self):
|
||||||
|
from sglang.srt.multimodal.processors.ernie45_vl import (
|
||||||
|
Ernie4_5_VLImageProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
config = {"image": {"max_pixels": 2000000}, "video": {"max_pixels": 500000}}
|
||||||
|
proc, mock_proc = self._make_override_processor(
|
||||||
|
Ernie4_5_VLImageProcessor, config
|
||||||
|
)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", images=["img1"], videos=["vid1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("images_kwargs"), {"max_pixels": 2000000}
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
call_kwargs.kwargs.get("videos_kwargs"), {"max_pixels": 500000}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_midashenglm_injects_audio_kwargs(self):
|
||||||
|
from sglang.srt.multimodal.processors.midashenglm import (
|
||||||
|
MiDashengLMMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
config = {"audio": {"sample_rate": 16000}}
|
||||||
|
proc, mock_proc = self._make_override_processor(
|
||||||
|
MiDashengLMMultimodalProcessor, config
|
||||||
|
)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", audios=["aud1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
audio_kw = call_kwargs.kwargs.get("audio_kwargs", {})
|
||||||
|
self.assertFalse(audio_kw.get("truncation", True))
|
||||||
|
self.assertEqual(audio_kw.get("sample_rate"), 16000)
|
||||||
|
|
||||||
|
def test_midashenglm_user_config_overrides_truncation(self):
|
||||||
|
"""User config can override the default truncation=False."""
|
||||||
|
from sglang.srt.multimodal.processors.midashenglm import (
|
||||||
|
MiDashengLMMultimodalProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
config = {"audio": {"truncation": True}}
|
||||||
|
proc, mock_proc = self._make_override_processor(
|
||||||
|
MiDashengLMMultimodalProcessor, config
|
||||||
|
)
|
||||||
|
|
||||||
|
proc.process_mm_data("test", audios=["aud1"])
|
||||||
|
|
||||||
|
call_kwargs = mock_proc.__call__.call_args
|
||||||
|
audio_kw = call_kwargs.kwargs.get("audio_kwargs", {})
|
||||||
|
# User config can override truncation if they explicitly set it
|
||||||
|
self.assertTrue(audio_kw.get("truncation"))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user