refactor: replace mm_inputs dict with MultimodalProcessorOutput (#21738)

This commit is contained in:
Mick
2026-04-03 23:26:37 +08:00
committed by GitHub
parent 9f409d0749
commit 030fb1c4b1
40 changed files with 408 additions and 314 deletions
@@ -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()
@@ -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:
+1 -1
View File
@@ -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
+3 -4
View File
@@ -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()
+58 -4
View File
@@ -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
+7 -7
View File
@@ -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)."""
@@ -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(
@@ -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."""
@@ -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
@@ -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(),
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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
@@ -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,
)
@@ -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
@@ -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):
@@ -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
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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,
)
@@ -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
@@ -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,
)
],
}
)