From 030fb1c4b10088f849d06fcd689c91cb62146338 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 3 Apr 2026 23:26:37 +0800 Subject: [PATCH] refactor: replace mm_inputs dict with MultimodalProcessorOutput (#21738) --- .../srt/disaggregation/encode_receiver.py | 2 +- .../srt/disaggregation/encode_server.py | 4 +- python/sglang/srt/managers/io_struct.py | 2 +- python/sglang/srt/managers/mm_utils.py | 7 +-- python/sglang/srt/managers/schedule_batch.py | 62 +++++++++++++++++-- python/sglang/srt/managers/scheduler.py | 14 ++--- .../sglang/srt/managers/session_controller.py | 2 +- .../sglang/srt/managers/tokenizer_manager.py | 18 +++--- .../multimodal/processors/base_processor.py | 17 ++--- .../sglang/srt/multimodal/processors/clip.py | 9 +-- .../srt/multimodal/processors/deepseek_ocr.py | 11 ++-- .../multimodal/processors/deepseek_vl_v2.py | 11 ++-- .../srt/multimodal/processors/dots_vlm.py | 15 ++--- .../srt/multimodal/processors/ernie45_vl.py | 23 ++++--- .../srt/multimodal/processors/gemma3.py | 13 ++-- .../srt/multimodal/processors/gemma3n.py | 14 ++--- .../sglang/srt/multimodal/processors/glm4v.py | 19 +++--- .../srt/multimodal/processors/glmasr.py | 15 ++--- .../srt/multimodal/processors/interns1pro.py | 42 +++++++------ .../srt/multimodal/processors/internvl.py | 49 ++++++++------- .../srt/multimodal/processors/janus_pro.py | 15 ++--- .../srt/multimodal/processors/kimi_k25.py | 15 +++-- .../srt/multimodal/processors/kimi_vl.py | 11 ++-- .../srt/multimodal/processors/lightonocr.py | 6 +- .../sglang/srt/multimodal/processors/llava.py | 14 +++-- .../srt/multimodal/processors/midashenglm.py | 18 +++--- .../srt/multimodal/processors/minicpm.py | 45 +++++++------- .../sglang/srt/multimodal/processors/mlama.py | 11 ++-- .../srt/multimodal/processors/mllama4.py | 15 ++--- .../multimodal/processors/nano_nemotron_vl.py | 17 ++--- .../sglang/srt/multimodal/processors/nvila.py | 13 ++-- .../srt/multimodal/processors/phi4mm.py | 13 ++-- .../srt/multimodal/processors/pixtral.py | 13 ++-- .../multimodal/processors/points_v15_chat.py | 11 ++-- .../srt/multimodal/processors/qwen_audio.py | 34 +++++----- .../srt/multimodal/processors/qwen_vl.py | 50 ++++++++------- .../processors/sarashina2_vision.py | 15 ++--- .../srt/multimodal/processors/step3_vl.py | 11 ++-- .../processors/transformers_auto.py | 32 +++++----- .../srt/multimodal/processors/whisper.py | 14 +++-- 40 files changed, 408 insertions(+), 314 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index 391bab6ef..5136e7865 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -532,7 +532,7 @@ class WaitingImageRequest: **self.recv_embedding_data.get_mm_extra_meta(), ) self.recv_req.mm_inputs = mm_inputs - self.recv_req.input_ids = mm_inputs["input_ids"] + self.recv_req.input_ids = mm_inputs.input_ids self.status = WaitingImageRequestStatus.SUCCESS self.recv_socket.close() diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 02d1054b4..b425d8c06 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -673,6 +673,7 @@ class MMEncoder: part_idx: int, hashes: Optional[List[str]] = None, ) -> torch.Tensor: + # mm_inputs: dict mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, modality) grid_thw = _get_mm_grid_dim(mm_inputs, modality) mm_feature = _convert(_get_mm_feature(mm_inputs, modality)) @@ -853,7 +854,6 @@ class MMEncoder: images = await self._flatten_and_load_images(mm_items) image_config = self.vision_config.get("image", {}) processor_input = self.image_processor(images=images, **image_config) - feature = processor_input["pixel_values"] if hasattr(self.model, "thinker"): # for omni models get_feature_method = self.model.thinker.get_image_feature else: @@ -908,7 +908,6 @@ class MMEncoder: ) processor_input["second_per_grid_ts"] = second_per_grid_ts_tensor - feature = processor_input["pixel_values_videos"] if hasattr(self.model, "thinker"): # for omni models get_feature_method = self.model.thinker.get_video_feature else: @@ -929,7 +928,6 @@ class MMEncoder: processor_input["audio_feature_lens_raw"] = input_lengths output_lengths = self._get_feat_extract_output_lengths(input_lengths) processor_input["audio_feature_lens"] = output_lengths - feature = processor_input["input_features"] if hasattr(self.model, "thinker"): # for omni models get_feature_method = self.model.thinker.get_audio_feature else: diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 471c01685..53a6fc902 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -674,7 +674,7 @@ class TokenizedGenerateReqInput(BaseReq): # The input token ids input_ids: List[int] # The multimodal inputs - mm_inputs: dict + mm_inputs: object # The sampling parameters sampling_params: SamplingParams # Whether to return the logprobs diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index f80d4a9f0..4555d3c2f 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1746,8 +1746,7 @@ def wrap_shm_features(obj): return obj if hasattr(obj, "mm_inputs") and obj.mm_inputs: - mm_items = obj.mm_inputs.get("mm_items", []) - for item in mm_items: + for item in obj.mm_inputs.mm_items: if ( hasattr(item, "feature") and isinstance(item.feature, torch.Tensor) @@ -1764,7 +1763,7 @@ def has_shm_features(recv_reqs): if has_shm_features(req.batch): return True elif hasattr(req, "mm_inputs") and req.mm_inputs: - for item in req.mm_inputs.get("mm_items", []): + for item in req.mm_inputs.mm_items: if isinstance(item.feature, ShmPointerMMData): return True return False @@ -1784,7 +1783,7 @@ def unwrap_shm_features(obj): return obj # Handle single requests if hasattr(obj, "mm_inputs") and obj.mm_inputs: - mm_items = obj.mm_inputs.get("mm_items", []) + mm_items = obj.mm_inputs.mm_items for item in mm_items: if isinstance(item.feature, ShmPointerMMData): item.feature = item.feature.materialize() diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 7cfa7d50a..b44c75a5d 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -353,6 +353,59 @@ class MultimodalDataItem: self.model_specific_data[extra_key] = extra_data +@dataclasses.dataclass +class MultimodalProcessorOutput: + """Raw output from multimodal processors, before pad/hash computation. + + This is the typed replacement for the dict previously returned by + ``BaseMultimodalProcessor.process_mm_data_async``. Unlike + ``MultimodalInputs``, items here do NOT carry pad_value or hash yet. + """ + + mm_items: List[MultimodalDataItem] + input_ids: Optional[List[int]] = None + + # image + im_token_id: Optional[int] = None + im_start_id: Optional[int] = None + im_end_id: Optional[int] = None + slice_start_id: Optional[int] = None + slice_end_id: Optional[int] = None + + # video + video_token_id: Optional[int] = None + + # audio + audio_token_id: Optional[int] = None + audio_start_id: Optional[int] = None + audio_end_id: Optional[int] = None + + # QWen2-VL related + mrope_positions: Optional[torch.Tensor] = None + mrope_position_delta: Optional[torch.Tensor] = None + + # for transformers-compatibility + token_type_ids: Optional[torch.Tensor] = None + + @staticmethod + def from_dict(d: dict) -> "MultimodalProcessorOutput": + return MultimodalProcessorOutput( + mm_items=d["mm_items"], + input_ids=d.get("input_ids"), + im_token_id=d.get("im_token_id"), + im_start_id=d.get("im_start_id"), + im_end_id=d.get("im_end_id"), + slice_start_id=d.get("slice_start_id"), + slice_end_id=d.get("slice_end_id"), + video_token_id=d.get("video_token_id"), + audio_token_id=d.get("audio_token_id"), + audio_start_id=d.get("audio_start_id"), + audio_end_id=d.get("audio_end_id"), + mrope_positions=d.get("mrope_positions"), + mrope_position_delta=d.get("mrope_position_delta"), + ) + + @dataclasses.dataclass class MultimodalInputs: """The multimodal data related inputs.""" @@ -388,8 +441,8 @@ class MultimodalInputs: item.feature = None @staticmethod - def from_dict(obj: dict): - mm_items = obj["mm_items"] + def from_processor_output(obj: "MultimodalProcessorOutput"): + mm_items = obj.mm_items for mm_item in mm_items: mm_item.reconstruct() @@ -442,8 +495,9 @@ class MultimodalInputs: "audio_token_id", ] for arg in optional_args: - if arg in obj: - setattr(ret, arg, obj[arg]) + val = getattr(obj, arg, None) + if val is not None: + setattr(ret, arg, val) return ret diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 58c56a516..ca4a7121d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1618,17 +1618,17 @@ class Scheduler( def _process_and_broadcast_mm_inputs( self, - raw_mm_inputs: Optional[dict], + raw_mm_inputs, ): """Materialize MultimodalInputs once on the entry rank and broadcast to others. Entry rank: - - constructs MultimodalInputs.from_dict(raw_mm_inputs) once + - constructs MultimodalInputs.from_processor_output() once - broadcasts to other ranks in self.cpu_group (if world_size > 1) Non-entry ranks: - receive the object via broadcast (if world_size > 1) - - otherwise (single-rank / no group) fall back to local from_dict + - otherwise (single-rank / no group) fall back to local from_processor_output Returns: MultimodalInputs | None @@ -1659,7 +1659,7 @@ class Scheduler( # increase the CUDA kernel launch time. if self.dp_tp_group.rank_in_group == 0: # Only the entry rank materializes once from dict. - image_inputs = MultimodalInputs.from_dict(raw_mm_inputs) + image_inputs = MultimodalInputs.from_processor_output(raw_mm_inputs) # Broadcast to other TP ranks (use src=0 within the group). if group_world_size > 1: obj_list = [image_inputs] @@ -1680,15 +1680,15 @@ class Scheduler( ) image_inputs = obj_list[0] else: - image_inputs = MultimodalInputs.from_dict(raw_mm_inputs) + image_inputs = MultimodalInputs.from_processor_output(raw_mm_inputs) return image_inputs - def _get_multimodal_inputs(self, mm_inputs_dict: dict): + def _get_multimodal_inputs(self, mm_inputs_dict): if self.server_args.enable_broadcast_mm_inputs_process: return self._process_and_broadcast_mm_inputs(mm_inputs_dict) else: - return MultimodalInputs.from_dict(mm_inputs_dict) + return MultimodalInputs.from_processor_output(mm_inputs_dict) def _maybe_compute_mrope_positions(self, req) -> None: """Compute M-RoPE positions when they are missing (e.g. gRPC preprocessed path).""" diff --git a/python/sglang/srt/managers/session_controller.py b/python/sglang/srt/managers/session_controller.py index 836feacb2..caf165c32 100644 --- a/python/sglang/srt/managers/session_controller.py +++ b/python/sglang/srt/managers/session_controller.py @@ -167,7 +167,7 @@ class Session: # Adjust mm_item offsets since they were computed on # the pre-strip sequence (with BOS at position 0) if req.mm_inputs: - for item in req.mm_inputs.get("mm_items", []): + for item in req.mm_inputs.mm_items: if item.offsets: if any(s == 0 for s, _ in item.offsets): logging.warning( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 63714087a..0d345ac6b 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -727,7 +727,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, ) if mm_inputs is None: - mm_inputs: Dict = await self.mm_processor.process_mm_data_async( + mm_inputs = await self.mm_processor.process_mm_data_async( image_data=obj.image_data, audio_data=obj.audio_data, input_text=(input_text or input_ids), @@ -741,7 +741,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi ): # In language_only mode with zmq_to_scheduler, if we didn't dispatch # to encoder (e.g., only one image), process locally like non-language_only mode - mm_inputs: Dict = await self.mm_processor.process_mm_data_async( + mm_inputs = await self.mm_processor.process_mm_data_async( image_data=obj.image_data, audio_data=obj.audio_data, input_text=(input_text or input_ids), @@ -749,18 +749,18 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi max_req_input_len=self.max_req_input_len, ) - if mm_inputs and "input_ids" in mm_inputs: - input_ids = mm_inputs["input_ids"] - if mm_inputs and "token_type_ids" in mm_inputs: - token_type_ids = mm_inputs.pop("token_type_ids") + if mm_inputs and mm_inputs.input_ids is not None: + input_ids = mm_inputs.input_ids + if mm_inputs and mm_inputs.token_type_ids is not None: + token_type_ids = mm_inputs.token_type_ids if not isinstance(token_type_ids, list): token_type_ids = token_type_ids.flatten().tolist() if ( envs.SGLANG_MM_PRECOMPUTE_HASH.get() and mm_inputs - and "mm_items" in mm_inputs + and mm_inputs.mm_items ): - for item in mm_inputs["mm_items"]: + for item in mm_inputs.mm_items: if isinstance(item, MultimodalDataItem): item.set_pad_value() else: @@ -931,7 +931,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi input_text: str, input_ids: List[int], input_embeds: Optional[Union[List[float], None]] = None, - mm_inputs: Optional[Dict] = None, + mm_inputs=None, token_type_ids: Optional[List[int]] = None, ) -> Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]: """Create a tokenized request object from common parameters.""" diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 773ce8620..839d5b74e 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -16,6 +16,7 @@ from sglang.srt.managers.schedule_batch import ( Modality, MultimodalDataItem, MultimodalInputFormat, + MultimodalProcessorOutput, ) from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( @@ -363,14 +364,14 @@ class BaseMultimodalProcessor(ABC): ) ) - return { - "input_ids": input_ids, - "mm_items": mm_items, - "im_start_id": self.IM_START_TOKEN_ID, - "im_end_id": self.IM_END_TOKEN_ID, - "im_token_id": self.IM_TOKEN_ID, - "video_token_id": getattr(self, "VIDEO_TOKEN_ID", None), - } + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=mm_items, + im_start_id=self.IM_START_TOKEN_ID, + im_end_id=self.IM_END_TOKEN_ID, + im_token_id=self.IM_TOKEN_ID, + video_token_id=getattr(self, "VIDEO_TOKEN_ID", None), + ) def process_mm_data( self, input_text, images=None, videos=None, audios=None, **kwargs diff --git a/python/sglang/srt/multimodal/processors/clip.py b/python/sglang/srt/multimodal/processors/clip.py index 19ff71e78..06f785b85 100644 --- a/python/sglang/srt/multimodal/processors/clip.py +++ b/python/sglang/srt/multimodal/processors/clip.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.clip import CLIPModel from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -29,7 +30,7 @@ class ClipImageProcessor(BaseMultimodalProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + ) diff --git a/python/sglang/srt/multimodal/processors/deepseek_ocr.py b/python/sglang/srt/multimodal/processors/deepseek_ocr.py index 9b9002d8d..becb0b2b3 100644 --- a/python/sglang/srt/multimodal/processors/deepseek_ocr.py +++ b/python/sglang/srt/multimodal/processors/deepseek_ocr.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.deepseek_ocr import DeepseekOCRForCausalLM from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -38,8 +39,8 @@ class DeepseekOCRProcessor(BaseMultimodalProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/deepseek_vl_v2.py b/python/sglang/srt/multimodal/processors/deepseek_vl_v2.py index 26708e8dc..3a9edd0e4 100644 --- a/python/sglang/srt/multimodal/processors/deepseek_vl_v2.py +++ b/python/sglang/srt/multimodal/processors/deepseek_vl_v2.py @@ -18,6 +18,7 @@ # CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.deepseek_vl2 import DeepseekVL2ForCausalLM from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -55,8 +56,8 @@ class DeepseekVL2ImageProcessor(BaseMultimodalProcessor): conversations=base_output.input_text, ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "im_token_id": self._processor.image_token_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_token_id=self._processor.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/dots_vlm.py b/python/sglang/srt/multimodal/processors/dots_vlm.py index 8d6faf5e8..c8e76562a 100644 --- a/python/sglang/srt/multimodal/processors/dots_vlm.py +++ b/python/sglang/srt/multimodal/processors/dots_vlm.py @@ -1,6 +1,7 @@ import re from typing import Dict, List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.dots_ocr import DotsOCRForCausalLM from sglang.srt.models.dots_vlm import DotsVLMForCausalLM from sglang.srt.multimodal.processors.base_processor import ( @@ -72,10 +73,10 @@ class DotsVLMImageProcessor(BaseMultimodalProcessor): if combined_mm_item is None: return None - return { - "input_ids": input_ids.tolist(), - "mm_items": combined_mm_item, - "im_start_id": self.im_start_id, - "im_end_id": self.im_end_id, - "im_token_id": self.image_token_id, - } + return MultimodalProcessorOutput( + mm_items=combined_mm_item, + input_ids=input_ids.tolist(), + im_token_id=self.image_token_id, + im_start_id=self.im_start_id, + im_end_id=self.im_end_id, + ) diff --git a/python/sglang/srt/multimodal/processors/ernie45_vl.py b/python/sglang/srt/multimodal/processors/ernie45_vl.py index b96ffb249..0607e600d 100644 --- a/python/sglang/srt/multimodal/processors/ernie45_vl.py +++ b/python/sglang/srt/multimodal/processors/ernie45_vl.py @@ -11,6 +11,7 @@ from transformers import BaseImageProcessorFast from sglang.srt.environ import envs from sglang.srt.layers.rotary_embedding import MRotaryEmbedding +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.ernie45_vl import Ernie4_5_VLMoeForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, @@ -423,15 +424,13 @@ class Ernie4_5_VLImageProcessor(SGLangBaseProcessor): input_ids.shape[0] == mrope_positions.shape[-1] ), "input_ids and mrope_positions should have the same length" - mm_inputs = { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_start_id": self.image_start_token_id, - "im_end_id": self.image_end_token_id, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "mrope_positions": mrope_positions, - "mrope_position_delta": mrope_position_delta, - } - - return mm_inputs + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_start_id=self.image_start_token_id, + im_end_id=self.image_end_token_id, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + mrope_positions=mrope_positions, + mrope_position_delta=mrope_position_delta, + ) diff --git a/python/sglang/srt/multimodal/processors/gemma3.py b/python/sglang/srt/multimodal/processors/gemma3.py index cbfb45e84..c6b35e843 100644 --- a/python/sglang/srt/multimodal/processors/gemma3.py +++ b/python/sglang/srt/multimodal/processors/gemma3.py @@ -4,6 +4,7 @@ from typing import Dict, List, Union from sglang.srt.managers.multimodal_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, ) +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.gemma3_mm import Gemma3ForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import MultimodalSpecialTokens @@ -46,9 +47,9 @@ class Gemma3SGLangImageProcessor(SGLangBaseProcessor): mm_items, input_ids, _ = self.process_and_combine_mm_data( base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_start_id": self.IM_START_TOKEN_ID, - "im_end_id": self.IM_END_TOKEN_ID, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_start_id=self.IM_START_TOKEN_ID, + im_end_id=self.IM_END_TOKEN_ID, + ) diff --git a/python/sglang/srt/multimodal/processors/gemma3n.py b/python/sglang/srt/multimodal/processors/gemma3n.py index 9ea8b8be3..5cb6d7962 100644 --- a/python/sglang/srt/multimodal/processors/gemma3n.py +++ b/python/sglang/srt/multimodal/processors/gemma3n.py @@ -17,6 +17,7 @@ from typing import Dict, List, Optional, Union from sglang.srt.managers.multimodal_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, ) +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.gemma3n_mm import Gemma3nForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import MultimodalSpecialTokens @@ -62,10 +63,9 @@ class Gemma3nSGLangProcessor(SGLangBaseProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - # TODO(mick): could we return MultimodalSpecialTokens directly? - "im_token_id": self.mm_tokens.image_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/glm4v.py b/python/sglang/srt/multimodal/processors/glm4v.py index 8169a764c..a44f14b6c 100644 --- a/python/sglang/srt/multimodal/processors/glm4v.py +++ b/python/sglang/srt/multimodal/processors/glm4v.py @@ -1,6 +1,7 @@ from typing import List, Union from sglang.srt.layers.rotary_embedding import MRotaryEmbedding +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.glm4v import Glm4vForConditionalGeneration from sglang.srt.models.glm4v_moe import Glm4vMoeForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( @@ -112,13 +113,11 @@ class Glm4vImageProcessor(SGLangBaseProcessor): ) mrope_positions = mrope_positions.squeeze(1) - mm_inputs = { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "mrope_positions": mrope_positions, - "mrope_position_delta": mrope_position_delta, - } - - return mm_inputs + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + mrope_positions=mrope_positions, + mrope_position_delta=mrope_position_delta, + ) diff --git a/python/sglang/srt/multimodal/processors/glmasr.py b/python/sglang/srt/multimodal/processors/glmasr.py index cebeb1f6a..1fcaf490a 100644 --- a/python/sglang/srt/multimodal/processors/glmasr.py +++ b/python/sglang/srt/multimodal/processors/glmasr.py @@ -1,5 +1,6 @@ import re +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.glmasr import GlmAsrForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -44,10 +45,10 @@ class GlmAsrProcessor(BaseMultimodalProcessor): mm_items, input_ids, ret = self.process_and_combine_mm_data( base_output, self.mm_tokens ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "audio_start_id": self.audio_start_id, - "audio_token_id": self.audio_token_id, - "audio_end_id": self.audio_end_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + audio_start_id=self.audio_start_id, + audio_token_id=self.audio_token_id, + audio_end_id=self.audio_end_id, + ) diff --git a/python/sglang/srt/multimodal/processors/interns1pro.py b/python/sglang/srt/multimodal/processors/interns1pro.py index d448dc1cb..0f4a909ad 100644 --- a/python/sglang/srt/multimodal/processors/interns1pro.py +++ b/python/sglang/srt/multimodal/processors/interns1pro.py @@ -1,7 +1,11 @@ import time from typing import List, Union -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.interns1pro import InternS1ProForConditionalGeneration from sglang.srt.multimodal.processors.qwen_vl import ( QwenVLImageProcessor, @@ -26,15 +30,15 @@ class InternS1_1ImageProcessor(QwenVLImageProcessor): ) ] - return { - "input_ids": input_ids, - "mm_items": mm_items, - "im_start_id": self.IM_START_TOKEN_ID, - "im_end_id": self.IM_END_TOKEN_ID, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=mm_items, + im_start_id=self.IM_START_TOKEN_ID, + im_end_id=self.IM_END_TOKEN_ID, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + ) async def process_mm_data_async( self, @@ -107,12 +111,12 @@ class InternS1_1ImageProcessor(QwenVLImageProcessor): f"total_time: {(get_rope_index_time - entry_time) * 1000:.2f} ms" ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_start_id": self.vision_start_token_id, - "im_end_id": self.vision_end_token_id, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_start_id=self.vision_start_token_id, + im_end_id=self.vision_end_token_id, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/internvl.py b/python/sglang/srt/multimodal/processors/internvl.py index c95d495ad..800f68110 100644 --- a/python/sglang/srt/multimodal/processors/internvl.py +++ b/python/sglang/srt/multimodal/processors/internvl.py @@ -11,6 +11,7 @@ from PIL import Image from sglang.srt.managers.schedule_batch import ( Modality, MultimodalDataItem, + MultimodalProcessorOutput, ) from sglang.srt.models.interns1 import InternS1ForConditionalGeneration from sglang.srt.models.internvl import InternVLChatModel @@ -337,14 +338,14 @@ class InternVLProcessor(BaseMultimodalProcessor): mm_token_id=mm_token_id, ) - return { - "input_ids": input_ids_tensor.flatten().tolist(), - "mm_items": mm_items, - "im_start_id": self.img_start_token_id, - "im_end_id": self.img_end_token_id, - "im_token_id": self.img_context_token_id, - "video_token_id": self.video_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids_tensor.flatten().tolist(), + mm_items=mm_items, + im_start_id=self.img_start_token_id, + im_end_id=self.img_end_token_id, + im_token_id=self.img_context_token_id, + video_token_id=self.video_token_id, + ) async def process_mm_data_async( self, image_data, input_text, request_obj, **kwargs @@ -610,14 +611,14 @@ class InternVLProcessor(BaseMultimodalProcessor): ) ) - return { - "input_ids": input_ids, - "mm_items": items, - "im_start_id": self.img_start_token_id, - "im_end_id": self.img_end_token_id, - "im_token_id": self.img_context_token_id, - "video_token_id": self.video_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=items, + im_start_id=self.img_start_token_id, + im_end_id=self.img_end_token_id, + im_token_id=self.img_context_token_id, + video_token_id=self.video_token_id, + ) async def process_internlm2_mm_data_async( self, image_data, input_text, request_obj, **kwargs @@ -728,11 +729,11 @@ class InternVLProcessor(BaseMultimodalProcessor): ) cumulative += num_patches - return { - "input_ids": input_ids, - "mm_items": items, - "im_start_id": self.img_start_token_id, - "im_end_id": self.img_end_token_id, - "im_token_id": self.img_context_token_id, - "video_token_id": self.video_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=items, + im_start_id=self.img_start_token_id, + im_end_id=self.img_end_token_id, + im_token_id=self.img_context_token_id, + video_token_id=self.video_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/janus_pro.py b/python/sglang/srt/multimodal/processors/janus_pro.py index 044e31dd2..f6711058d 100644 --- a/python/sglang/srt/multimodal/processors/janus_pro.py +++ b/python/sglang/srt/multimodal/processors/janus_pro.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.deepseek_janus_pro import MultiModalityCausalLM from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -35,10 +36,10 @@ class JanusProImageProcessor(BaseMultimodalProcessor): base_out, self.mm_tokens, prompt=base_out.input_text ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "im_start_id": self._processor.image_start_id, - "im_end_id": self._processor.image_end_id, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_start_id=self._processor.image_start_id, + im_end_id=self._processor.image_end_id, + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/kimi_k25.py b/python/sglang/srt/multimodal/processors/kimi_k25.py index cef3e6933..ece7b02c3 100644 --- a/python/sglang/srt/multimodal/processors/kimi_k25.py +++ b/python/sglang/srt/multimodal/processors/kimi_k25.py @@ -3,7 +3,10 @@ from typing import Dict, List, Tuple, Union import torch -from sglang.srt.managers.schedule_batch import MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.kimi_k25 import KimiK25ForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, @@ -46,11 +49,11 @@ class KimiK2_5VLImageProcessor(SGLangBaseProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + ) def _process_and_collect_mm_items( self, input_text: str, images=None, audios=None, videos=None, **kwargs diff --git a/python/sglang/srt/multimodal/processors/kimi_vl.py b/python/sglang/srt/multimodal/processors/kimi_vl.py index b466f1b40..e98ec0cd3 100644 --- a/python/sglang/srt/multimodal/processors/kimi_vl.py +++ b/python/sglang/srt/multimodal/processors/kimi_vl.py @@ -1,6 +1,7 @@ import re from typing import Dict, List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.kimi_vl import KimiVLForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, @@ -42,8 +43,8 @@ class KimiVLImageProcessor(SGLangBaseProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/lightonocr.py b/python/sglang/srt/multimodal/processors/lightonocr.py index 59c6c1429..cc687e5c8 100644 --- a/python/sglang/srt/multimodal/processors/lightonocr.py +++ b/python/sglang/srt/multimodal/processors/lightonocr.py @@ -80,8 +80,8 @@ class LightOnOCRProcessor(PixtralProcessor): return result # Remove break/end tokens and fix multimodal item offsets - input_ids = result.get("input_ids", []) - mm_items = result.get("mm_items", []) + input_ids = result.input_ids or [] + mm_items = result.mm_items or [] new_input_ids = [] old_to_new = {} @@ -106,5 +106,5 @@ class LightOnOCRProcessor(PixtralProcessor): if new_indices: mm_item.offsets = [(new_indices[0], new_indices[-1])] - result["input_ids"] = new_input_ids + result.input_ids = new_input_ids return result diff --git a/python/sglang/srt/multimodal/processors/llava.py b/python/sglang/srt/multimodal/processors/llava.py index 8729e8547..bbfd41016 100644 --- a/python/sglang/srt/multimodal/processors/llava.py +++ b/python/sglang/srt/multimodal/processors/llava.py @@ -8,7 +8,11 @@ from transformers.models.auto.processing_auto import ( ) import sglang.srt.managers.multimodal_processor as sgl_mm_processor_utils -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.llava import ( LlavaForConditionalGeneration, LlavaLlamaForCausalLM, @@ -139,7 +143,7 @@ class LlavaImageProcessor(BaseMultimodalProcessor): model_specific_data=item, ) ) - return {"mm_items": mm_items} + return MultimodalProcessorOutput(mm_items=mm_items) async def process_mm_data_async( self, @@ -218,9 +222,9 @@ class LlavaImageProcessor(BaseMultimodalProcessor): ) ) - return { - "mm_items": mm_items, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + ) class LlavaMultimodalProcessor(BaseMultimodalProcessor): diff --git a/python/sglang/srt/multimodal/processors/midashenglm.py b/python/sglang/srt/multimodal/processors/midashenglm.py index 570765b68..2aaa31214 100644 --- a/python/sglang/srt/multimodal/processors/midashenglm.py +++ b/python/sglang/srt/multimodal/processors/midashenglm.py @@ -3,7 +3,7 @@ import re import torch -from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.managers.schedule_batch import Modality, MultimodalProcessorOutput from sglang.srt.models.midashenglm import MiDashengLMModel from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -154,12 +154,12 @@ class MiDashengLMMultimodalProcessor(BaseMultimodalProcessor): mm_items[0].audio_length = audio_length logger.info(f"Set audio_length={audio_length} (fallback, waveform length)") - result = { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "audio_start_id": self.audio_start_id, - "audio_token_id": self.audio_token_id, - "audio_end_id": self.audio_end_id, - } - logger.info(f"Returning {len(result['mm_items'])} mm_items") + result = MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + audio_start_id=self.audio_start_id, + audio_token_id=self.audio_token_id, + audio_end_id=self.audio_end_id, + ) + logger.info(f"Returning {len(result.mm_items)} mm_items") return result diff --git a/python/sglang/srt/multimodal/processors/minicpm.py b/python/sglang/srt/multimodal/processors/minicpm.py index 613079e04..d4c407c13 100644 --- a/python/sglang/srt/multimodal/processors/minicpm.py +++ b/python/sglang/srt/multimodal/processors/minicpm.py @@ -5,6 +5,7 @@ import torch from sglang.srt.managers.schedule_batch import ( Modality, MultimodalDataItem, + MultimodalProcessorOutput, ) from sglang.srt.models.minicpmo import MiniCPMO from sglang.srt.models.minicpmv import MiniCPMV @@ -158,17 +159,17 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor): mm_end_id=self.audio_end_id, ) - return { - "mm_items": mm_items, - "input_ids": input_ids_tensor.flatten().tolist(), - "audio_start_id": self.audio_start_id, - "audio_end_id": self.audio_end_id, - "im_token_id": self.im_token_id, - "im_start_id": self.im_start_id, - "im_end_id": self.im_end_id, - "slice_start_id": self.slice_start_id, - "slice_end_id": self.slice_end_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids_tensor.flatten().tolist(), + audio_start_id=self.audio_start_id, + audio_end_id=self.audio_end_id, + im_token_id=self.im_token_id, + im_start_id=self.im_start_id, + im_end_id=self.im_end_id, + slice_start_id=self.slice_start_id, + slice_end_id=self.slice_end_id, + ) async def process_mm_data_async( self, @@ -291,14 +292,14 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor): modality=Modality.AUDIO, ) items += [item] - return { - "mm_items": items, - "input_ids": input_ids.tolist(), - "audio_start_id": self.audio_start_id, - "audio_end_id": self.audio_end_id, - "im_token_id": self.im_token_id, - "im_start_id": self.im_start_id, - "im_end_id": self.im_end_id, - "slice_start_id": self.slice_start_id, - "slice_end_id": self.slice_end_id, - } + return MultimodalProcessorOutput( + mm_items=items, + input_ids=input_ids.tolist(), + audio_start_id=self.audio_start_id, + audio_end_id=self.audio_end_id, + im_token_id=self.im_token_id, + im_start_id=self.im_start_id, + im_end_id=self.im_end_id, + slice_start_id=self.slice_start_id, + slice_end_id=self.slice_end_id, + ) diff --git a/python/sglang/srt/multimodal/processors/mlama.py b/python/sglang/srt/multimodal/processors/mlama.py index 432215a4f..52129765c 100644 --- a/python/sglang/srt/multimodal/processors/mlama.py +++ b/python/sglang/srt/multimodal/processors/mlama.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.mllama import MllamaForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -30,8 +31,8 @@ class MllamaImageProcessor(BaseMultimodalProcessor): base_out, self.mm_tokens ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/mllama4.py b/python/sglang/srt/multimodal/processors/mllama4.py index 4f04688b8..3983df275 100644 --- a/python/sglang/srt/multimodal/processors/mllama4.py +++ b/python/sglang/srt/multimodal/processors/mllama4.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.mllama4 import Llama4ForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -40,10 +41,10 @@ class Mllama4ImageProcessor(BaseMultimodalProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_start_id": self.IM_START_TOKEN_ID, - "im_end_id": self.IM_END_TOKEN_ID, - "im_token_id": self.IM_TOKEN_ID, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_start_id=self.IM_START_TOKEN_ID, + im_end_id=self.IM_END_TOKEN_ID, + im_token_id=self.IM_TOKEN_ID, + ) diff --git a/python/sglang/srt/multimodal/processors/nano_nemotron_vl.py b/python/sglang/srt/multimodal/processors/nano_nemotron_vl.py index 98986090f..90f283ae8 100644 --- a/python/sglang/srt/multimodal/processors/nano_nemotron_vl.py +++ b/python/sglang/srt/multimodal/processors/nano_nemotron_vl.py @@ -18,6 +18,7 @@ import torch from PIL import Image from sglang.srt.configs.nano_nemotron_vl import NemotronH_Nano_VL_V2_Config +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.nano_nemotron_vl import NemotronH_Nano_VL_V2 from sglang.srt.multimodal.evs import EVSProcessor from sglang.srt.multimodal.internvl_utils import image_to_pixel_values @@ -202,11 +203,11 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor): input_ids_list=prompt_ids_list, ) - return { - "input_ids": prompt_ids_list, - "mm_items": items, - "im_start_id": self.img_start_token_id, - "im_end_id": self.img_end_token_id, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + input_ids=prompt_ids_list, + mm_items=items, + im_start_id=self.img_start_token_id, + im_end_id=self.img_end_token_id, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/nvila.py b/python/sglang/srt/multimodal/processors/nvila.py index f34d600b3..5fe64d10d 100644 --- a/python/sglang/srt/multimodal/processors/nvila.py +++ b/python/sglang/srt/multimodal/processors/nvila.py @@ -6,6 +6,7 @@ from transformers.processing_utils import ProcessorMixin from transformers.tokenization_utils_base import PreTrainedTokenizerBase from sglang.srt.managers.io_struct import GenerateReqInput +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.jet_vlm import JetVLMForConditionalGeneration from sglang.srt.models.nvila import NVILAForConditionalGeneration from sglang.srt.models.nvila_lite import NVILALiteForConditionalGeneration @@ -71,9 +72,9 @@ class NVILAMultimodalProcessor(BaseMultimodalProcessor): num_frames=NUM_VIDEO_FRAMES, ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/phi4mm.py b/python/sglang/srt/multimodal/processors/phi4mm.py index c59a41685..6ae194eac 100644 --- a/python/sglang/srt/multimodal/processors/phi4mm.py +++ b/python/sglang/srt/multimodal/processors/phi4mm.py @@ -3,6 +3,7 @@ from typing import List, Union from transformers.processing_utils import ProcessorMixin +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.phi4mm import Phi4MMForCausalLM from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -92,9 +93,9 @@ class Phi4MMMultimodalProcessor(BaseMultimodalProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/pixtral.py b/python/sglang/srt/multimodal/processors/pixtral.py index ed40fc017..963bd6820 100644 --- a/python/sglang/srt/multimodal/processors/pixtral.py +++ b/python/sglang/srt/multimodal/processors/pixtral.py @@ -6,7 +6,7 @@ from transformers.models.pixtral.image_processing_pixtral import ( _num_image_tokens as _get_pixtral_hf_num_image_tokens, ) -from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.managers.schedule_batch import Modality, MultimodalProcessorOutput from sglang.srt.models.pixtral import ( PixtralForConditionalGeneration, PixtralVisionModel, @@ -127,9 +127,8 @@ class PixtralProcessor(BaseMultimodalProcessor): mm_data, self.mm_tokens ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "im_token_id": self.IM_TOKEN_ID, - "im_token": self.image_token, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_token_id=self.IM_TOKEN_ID, + ) diff --git a/python/sglang/srt/multimodal/processors/points_v15_chat.py b/python/sglang/srt/multimodal/processors/points_v15_chat.py index be23c28db..7fac7e909 100644 --- a/python/sglang/srt/multimodal/processors/points_v15_chat.py +++ b/python/sglang/srt/multimodal/processors/points_v15_chat.py @@ -2,6 +2,7 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.points_v15_chat import POINTSV15ChatModel from sglang.srt.multimodal.processors.qwen_vl import QwenVLImageProcessor @@ -35,8 +36,8 @@ class POINTSV15ChatProcessor(QwenVLImageProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/qwen_audio.py b/python/sglang/srt/multimodal/processors/qwen_audio.py index 90c2ffd45..5ca7c957c 100644 --- a/python/sglang/srt/multimodal/processors/qwen_audio.py +++ b/python/sglang/srt/multimodal/processors/qwen_audio.py @@ -1,6 +1,10 @@ import re -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.qwen2_audio import Qwen2AudioForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -69,13 +73,13 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor): if mm_items: mm_items[0].audio_feature_lens = output_lengths - return { - "mm_items": mm_items, - "input_ids": input_ids, - "audio_start_id": self.audio_start_id, - "audio_token_id": self.audio_token_id, - "audio_end_id": self.audio_end_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids, + audio_start_id=self.audio_start_id, + audio_token_id=self.audio_token_id, + audio_end_id=self.audio_end_id, + ) async def process_mm_data_async( self, @@ -104,10 +108,10 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor): mm_items[0].audio_feature_lens = output_lengths - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "audio_start_id": self.audio_start_id, - "audio_token_id": self.audio_token_id, - "audio_end_id": self.audio_end_id, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + audio_start_id=self.audio_start_id, + audio_token_id=self.audio_token_id, + audio_end_id=self.audio_end_id, + ) diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index 95c7cd21a..3f102567d 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -12,7 +12,11 @@ from torchvision.transforms import InterpolationMode from sglang.srt.environ import envs from sglang.srt.layers.rotary_embedding import MRotaryEmbedding -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLForConditionalGeneration from sglang.srt.models.qwen2_vl import Qwen2VLForConditionalGeneration from sglang.srt.models.qwen3_5 import ( @@ -474,17 +478,17 @@ class QwenVLImageProcessor(SGLangBaseProcessor): ) ) - return { - "input_ids": input_ids, - "mm_items": mm_items, - "im_start_id": self.IM_START_TOKEN_ID, - "im_end_id": self.IM_END_TOKEN_ID, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - "mrope_positions": mrope_positions, - "mrope_position_delta": mrope_position_delta, - } + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=mm_items, + im_start_id=self.IM_START_TOKEN_ID, + im_end_id=self.IM_END_TOKEN_ID, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + mrope_positions=mrope_positions, + mrope_position_delta=mrope_position_delta, + ) async def process_mm_data_async( self, @@ -599,14 +603,14 @@ class QwenVLImageProcessor(SGLangBaseProcessor): f"total_time: {(get_rope_index_time - entry_time) * 1000:.2f} ms" ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_start_id": self.vision_start_token_id, - "im_end_id": self.vision_end_token_id, - "im_token_id": self.mm_tokens.image_token_id, - "video_token_id": self.mm_tokens.video_token_id, - "audio_token_id": self.mm_tokens.audio_token_id, - "mrope_positions": mrope_positions, - "mrope_position_delta": mrope_position_delta, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_start_id=self.vision_start_token_id, + im_end_id=self.vision_end_token_id, + im_token_id=self.mm_tokens.image_token_id, + video_token_id=self.mm_tokens.video_token_id, + audio_token_id=self.mm_tokens.audio_token_id, + mrope_positions=mrope_positions, + mrope_position_delta=mrope_position_delta, + ) diff --git a/python/sglang/srt/multimodal/processors/sarashina2_vision.py b/python/sglang/srt/multimodal/processors/sarashina2_vision.py index fc7bdf3c9..c56f969c6 100644 --- a/python/sglang/srt/multimodal/processors/sarashina2_vision.py +++ b/python/sglang/srt/multimodal/processors/sarashina2_vision.py @@ -1,5 +1,6 @@ from typing import List, Union +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.sarashina2_vision import Sarashina2VisionForCausalLM from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, @@ -72,10 +73,10 @@ class Sarashina2VisionProcessor(BaseMultimodalProcessor): mm_tokens=self.mm_tokens, ) - return { - "mm_items": mm_items, - "input_ids": input_ids.tolist(), - "im_token_id": self.mm_tokens.image_token_id, - "im_start_id": self.IM_START_ID, - "im_end_id": self.IM_END_ID, - } + return MultimodalProcessorOutput( + mm_items=mm_items, + input_ids=input_ids.tolist(), + im_token_id=self.mm_tokens.image_token_id, + im_start_id=self.IM_START_ID, + im_end_id=self.IM_END_ID, + ) diff --git a/python/sglang/srt/multimodal/processors/step3_vl.py b/python/sglang/srt/multimodal/processors/step3_vl.py index d11d97787..e31985192 100644 --- a/python/sglang/srt/multimodal/processors/step3_vl.py +++ b/python/sglang/srt/multimodal/processors/step3_vl.py @@ -10,6 +10,7 @@ from torchvision import transforms from torchvision.transforms import InterpolationMode from transformers import BatchFeature, ProcessorMixin, TensorType +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.step3_vl import Step3VLForConditionalGeneration from sglang.srt.models.step3_vl_10b import StepVLForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( @@ -514,8 +515,8 @@ class Step3VLImageProcessor(SGLangBaseProcessor): base_output, self.mm_tokens ) - return { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - "im_token_id": self.mm_tokens.image_token_id, - } + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/multimodal/processors/transformers_auto.py b/python/sglang/srt/multimodal/processors/transformers_auto.py index b99f06616..579ae6e24 100644 --- a/python/sglang/srt/multimodal/processors/transformers_auto.py +++ b/python/sglang/srt/multimodal/processors/transformers_auto.py @@ -2,7 +2,11 @@ from typing import Optional import torch -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, MultimodalSpecialTokens, @@ -166,10 +170,10 @@ class TransformersAutoMultimodalProcessor(BaseMultimodalProcessor): # Build mm_items from processor output mm_items = self._build_mm_items(processor_output, input_ids) - ret = { - "input_ids": input_ids.tolist(), - "mm_items": mm_items, - } + ret = MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + ) # Propagate token_type_ids for models that need it (Gemma3, PaliGemma) token_type_key = ( @@ -178,14 +182,14 @@ class TransformersAutoMultimodalProcessor(BaseMultimodalProcessor): else "token_type_ids" ) if token_type_key in processor_output: - ret["token_type_ids"] = processor_output[token_type_key].flatten().tolist() + ret.token_type_ids = processor_output[token_type_key].flatten().tolist() if self.mm_tokens.image_token_id is not None: - ret["im_token_id"] = self.mm_tokens.image_token_id + ret.im_token_id = self.mm_tokens.image_token_id if self.mm_tokens.video_token_id is not None: - ret["video_token_id"] = self.mm_tokens.video_token_id + ret.video_token_id = self.mm_tokens.video_token_id if self.mm_tokens.audio_token_id is not None: - ret["audio_token_id"] = self.mm_tokens.audio_token_id + ret.audio_token_id = self.mm_tokens.audio_token_id image_start_id = _first_attr( self.hf_config, @@ -196,20 +200,20 @@ class TransformersAutoMultimodalProcessor(BaseMultimodalProcessor): ("image_end_token_id", "vision_end_token_id", "im_end_id"), ) if image_start_id is not None: - ret["im_start_id"] = image_start_id + ret.im_start_id = image_start_id if image_end_id is not None: - ret["im_end_id"] = image_end_id + ret.im_end_id = image_end_id # M-RoPE positions (Qwen2.5-VL, Qwen3-VL) if self._is_mrope: image_grid_thw = processor_output.get("image_grid_thw") video_grid_thw = processor_output.get("video_grid_thw") mrope_positions, mrope_position_delta = self._compute_mrope_positions( - ret["input_ids"], + ret.input_ids, image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, ) - ret["mrope_positions"] = mrope_positions - ret["mrope_position_delta"] = mrope_position_delta + ret.mrope_positions = mrope_positions + ret.mrope_position_delta = mrope_position_delta return ret diff --git a/python/sglang/srt/multimodal/processors/whisper.py b/python/sglang/srt/multimodal/processors/whisper.py index c09aa8854..a4472991e 100644 --- a/python/sglang/srt/multimodal/processors/whisper.py +++ b/python/sglang/srt/multimodal/processors/whisper.py @@ -1,7 +1,11 @@ import logging from typing import Any, Dict, Optional -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.whisper import WhisperForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor from sglang.srt.utils import load_audio @@ -187,12 +191,12 @@ class WhisperProcessor(BaseMultimodalProcessor): return_tensors="pt", )["input_features"][0] - return { - "input_ids": input_ids, - "mm_items": [ + return MultimodalProcessorOutput( + input_ids=input_ids, + mm_items=[ MultimodalDataItem( feature=input_features, modality=Modality.AUDIO, ) ], - } + )