Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
31 commits
Select commit Hold shift + click to select a range
554d3c2
refactor: simplify multimodal data processing
JustinTong0323 Jul 17, 2025
75dd9fd
revert modification for mllama4
JustinTong0323 Jul 17, 2025
ce0dceb
fix lint
JustinTong0323 Jul 17, 2025
077163b
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 18, 2025
a5b1108
refactor: update type hints for model_specific_data and items in sche…
JustinTong0323 Jul 18, 2025
4cdae8f
Refactors how model-specific multimodal data is handled.
JustinTong0323 Jul 19, 2025
88a27e3
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 19, 2025
b764985
fix: restore change fot models/mllama.py
JustinTong0323 Jul 19, 2025
607ec41
fix: unify multimodal token handling across processors
JustinTong0323 Jul 19, 2025
241a664
misc: enhance type hints in base_processor.py
JustinTong0323 Jul 19, 2025
4b31269
fix: use get method for model_specific_data access in multiple models
JustinTong0323 Jul 19, 2025
7bfb3f5
Removes `max_req_input_len` from multimodal processor calls.
JustinTong0323 Jul 19, 2025
2eacc3e
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 19, 2025
2bdbd62
fix after merge main:
JustinTong0323 Jul 19, 2025
4ed0be3
Refactors audio processing in Qwen2
JustinTong0323 Jul 19, 2025
3e405d7
Refactors multimodal token handling
JustinTong0323 Jul 19, 2025
13131ca
Rename `precomputed_features` to `precomputed_embeddings`
JustinTong0323 Jul 19, 2025
62c663d
Remove `image_sizes` as common MultimodalDataItem attributes
JustinTong0323 Jul 19, 2025
287baac
fix: deepseek vl related issue
JustinTong0323 Jul 19, 2025
d1d643b
Refactors MultimodalDataItem attr access and usage
JustinTong0323 Jul 19, 2025
840ddfe
misc: remove redundant attr in MllamaImageProcessor
JustinTong0323 Jul 19, 2025
3f63a14
fix: update attribute access in MllamaForConditionalGeneration
JustinTong0323 Jul 19, 2025
edc67d6
fix lint
JustinTong0323 Jul 19, 2025
d2ab9d4
fix llama atttr
JustinTong0323 Jul 19, 2025
a5c64ed
Temporarily disable MllamaServer tests due to instability
JustinTong0323 Jul 19, 2025
d8776da
fix: qwen vl video attr error
JustinTong0323 Jul 19, 2025
8360efd
fix lint
JustinTong0323 Jul 19, 2025
609c6f9
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 19, 2025
a8385dc
Streamlines access to `MultimodalDataItem` attributes by using `__get…
JustinTong0323 Jul 20, 2025
9691f9c
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 20, 2025
d106df8
Merge branch 'main' into feat-simlify-multimodal-dataitem-processing
JustinTong0323 Jul 20, 2025
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 15 additions & 26 deletions python/sglang/srt/managers/schedule_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@
import threading
from enum import Enum, auto
from http import HTTPStatus
from typing import TYPE_CHECKING, Any, List, Optional, Set, Tuple, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union

import numpy as np
import torch
Expand Down Expand Up @@ -201,7 +201,7 @@ class MultimodalDataItem:
For example, if there are 3 images and 1 audio inputs, there will be 2 MultimodalDataItem.
One for images and one for audio.

We put the common fields first and the model-specific fields last.
We put the common fields first and the model-specific fields in model_specific_data.
"""

modality: Modality
Expand All @@ -211,33 +211,22 @@ class MultimodalDataItem:
# the raw features returned by processor, e.g. pixel_values or audio_features
feature: Union[torch.Tensor, np.ndarray] = None

# Common fields used across multiple models
image_sizes: Tuple[int, int] = None

audio_feature_lens: Optional[List[torch.Tensor]] = None
audio_offsets: Optional[List[Tuple[int, int]]] = None
precomputed_features: Optional[Union[torch.Tensor, np.ndarray]] = None

# For qwen-vl
image_grid_thw: Union[torch.Tensor, np.ndarray] = None
second_per_grid_ts: Optional[List[torch.Tensor]] = None

# For deepseek-vl
image_emb_mask: Optional[torch.Tensor] = None
image_spatial_crop: Optional[torch.Tensor] = None

# For minicpmv
# [num_images, (n, w, h)]
tgt_size: Tuple[int, int] = None

# For mllama
aspect_ratio_id: Optional[List[torch.Tensor]] = None
aspect_ratio_mask: Optional[List[torch.Tensor]] = None

# For kimi-vl
image_grid_hws: Optional[List[torch.Tensor]] = None

# For gemma3n
input_features_mask: Optional[torch.Tensor] = None
# Model-specific data stored in a dictionary
# This should contains all the individual model-specific fields like:
# - image_grid_thw, second_per_grid_ts (qwen-vl)
# - image_emb_mask, image_spatial_crop (deepseek-vl)
# - tgt_size (minicpmv)
# - aspect_ratio_id, aspect_ratio_mask (mllama)
# - image_grid_hws (kimi-vl)
# - input_features_mask (gemma3n)
# - audio_feature_lens, audio_offsets (minicpmo)
# - ...
model_specific_data: Dict[str, Any] = dataclasses.field(default_factory=dict)

@staticmethod
def is_empty_list(l):
Expand Down Expand Up @@ -303,7 +292,7 @@ def from_dict(obj: dict):
def merge(self, other):
self.feature += other.feature
self.image_sizes += other.image_sizes
self.image_offsets += other.image_offsets
self.offsets += other.offsets
self.hash = hash((self.hash, other.hash))
self.set_pad_value()
Comment thread
JustinTong0323 marked this conversation as resolved.

Expand Down
13 changes: 2 additions & 11 deletions python/sglang/srt/models/mllama.py
Original file line number Diff line number Diff line change
Expand Up @@ -836,19 +836,10 @@ def __init__(
prefix="multi_modal_projector",
)
self.logits_processor = LogitsProcessor(config.text_config)
self.padding_pattern = MultiModalityDataPaddingPatternMultimodalTokens()

def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
Comment thread
mickqian marked this conversation as resolved.
pixel_values = torch.cat([item.feature for item in mm_inputs.mm_items], dim=0)
pad_values = [item.pad_value for item in mm_inputs.mm_items]

num_concurrent_media, num_tiles = pixel_values.shape[1:3]
num_patches = self.vision_model.num_patches
image_len = num_concurrent_media * num_tiles * num_patches
mm_inputs.num_image_tokens = image_len

pad_ids = pad_values * ((image_len + len(pad_values)) // len(pad_values))

return pad_ids[:image_len] + input_ids
return self.padding_pattern.pad_input_tokens(input_ids, mm_inputs)

def _batch_image_inputs(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode() or all(forward_batch.encoder_cached):
Expand Down
4 changes: 2 additions & 2 deletions python/sglang/srt/models/mllama4.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@ def __init__(
self.logits_processor = LogitsProcessor(
config.text_config if hasattr(config, "text_config") else config
)
self.padding_pattern = MultiModalityDataPaddingPatternMultimodalTokens()

def _has_vision_weights(self, config) -> bool:
"""Check if the model has vision components by examining the checkpoint."""
Expand Down Expand Up @@ -135,8 +136,7 @@ def _check_vision_weights_in_index(self, index_file: str) -> bool:
return False

def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
return pattern.pad_input_tokens(input_ids, mm_inputs)
return self.padding_pattern.pad_input_tokens(input_ids, mm_inputs)

def get_image_feature(
self,
Expand Down
6 changes: 3 additions & 3 deletions python/sglang/srt/multimodal/processors/base_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -529,9 +529,9 @@ def collect_mm_items_from_processor_output(

if attr_name in self.FEATURE_NAMES:
attr_name = "feature"

# Set attribute
setattr(items[modality], attr_name, value)
setattr(items[modality], attr_name, value)
Comment thread
JustinTong0323 marked this conversation as resolved.
Outdated
else:
items[modality].model_specific_data[attr_name] = value

return list(items.values())

Expand Down
37 changes: 18 additions & 19 deletions python/sglang/srt/multimodal/processors/clip.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
from typing import List, Union

from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.models.clip import CLIPModel
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.utils import load_image
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)


class ClipImageProcessor(BaseMultimodalProcessor):
Expand All @@ -15,19 +16,17 @@ def __init__(self, hf_config, server_args, _processor):
async def process_mm_data_async(
self, image_data: List[Union[str, bytes]], input_text, *args, **kwargs
):
if isinstance(input_text, list):
assert len(input_text) and isinstance(input_text[0], int)
input_text = self._processor.tokenizer.decode(input_text)

images = [load_image(image)[0] for image in image_data]

image_inputs = self.process_mm_data(input_text=input_text, images=images)
image_inputs["data_hashes"] = [hash(str(image_data))]
image_inputs["input_ids"] = image_inputs["input_ids"].tolist()[0]
image_inputs["mm_items"] = [
MultimodalDataItem(
feature=image_inputs["pixel_values"], modality=Modality.IMAGE
)
]

return image_inputs
base_output = self.load_mm_data(
prompt=input_text,
multimodal_tokens=MultimodalSpecialTokens(
image_token=self._processor.tokenizer.image_token
),
image_data=image_data,
)

mm_items, input_ids, _ = self.process_and_combine_mm_data(base_output)

return {
"input_ids": input_ids.tolist(),
"mm_items": mm_items,
}
32 changes: 2 additions & 30 deletions python/sglang/srt/multimodal/processors/deepseek_vl_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,9 +18,6 @@
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
from typing import List, Union

import torch

from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.models.deepseek_vl2 import DeepseekVL2ForCausalLM
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
Expand All @@ -39,7 +36,6 @@ async def process_mm_data_async(
self,
image_data: List[Union[str, bytes]],
input_text,
request_obj,
max_req_input_len,
*args,
**kwargs
Expand All @@ -50,34 +46,10 @@ async def process_mm_data_async(
multimodal_tokens=MultimodalSpecialTokens(image_token=self.IMAGE_TOKEN),
max_req_input_len=max_req_input_len,
)
res = self.process_mm_data(
input_text=base_output.input_text,
images=base_output.images,
max_req_input_len=max_req_input_len,
conversations=base_output.input_text,
)
images_seq_mask = res["images_seq_mask"]
images_spatial_crop = res["images_spatial_crop"]
batched_images_spatial_crop = []
batched_images_spatial_crop.append(images_spatial_crop)
batched_images_spatial_crop = torch.stack(batched_images_spatial_crop, dim=0)

items = []
input_ids = res["input_ids"]
image_offsets = self.get_mm_items_offset(
input_ids=input_ids, mm_token_id=self._processor.image_token_id
)
item = MultimodalDataItem(
feature=res["images"],
offsets=image_offsets,
modality=Modality.IMAGE,
image_emb_mask=images_seq_mask,
image_spatial_crop=batched_images_spatial_crop,
)
items += [item]
mm_items, input_ids, _ = self.process_and_combine_mm_data(base_output)

return {
"mm_items": items,
"mm_items": mm_items,
"input_ids": input_ids.tolist(),
"im_token_id": self._processor.image_token_id,
}
20 changes: 2 additions & 18 deletions python/sglang/srt/multimodal/processors/janus_pro.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,26 +33,10 @@ async def process_mm_data_async(
max_req_input_len=max_req_input_len,
)

images = base_out.images
res = self.process_mm_data(
input_text=base_out.input_text,
prompt=base_out.input_text,
images=images,
)
mm_items, input_ids, _ = self.process_and_combine_mm_data(base_out)

input_ids = res["input_ids"].flatten()
image_offsets = self.get_mm_items_offset(
input_ids=input_ids, mm_token_id=processor.image_id
)
return {
"mm_items": [
MultimodalDataItem(
feature=res["pixel_values"],
image_emb_mask=res["images_emb_mask"],
offsets=image_offsets,
modality=Modality.IMAGE,
)
],
"mm_items": mm_items,
"input_ids": input_ids.tolist(),
"im_start_id": processor.image_start_id,
"im_end_id": processor.image_end_id,
Expand Down
4 changes: 2 additions & 2 deletions python/sglang/srt/multimodal/processors/minicpm.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ async def process_mm_data_async(
item = MultimodalDataItem(
feature=pixel_values,
offsets=image_offsets,
tgt_size=tgt_sizes_flat,
model_specific_data={"tgt_size": tgt_sizes_flat},
modality=Modality.IMAGE,
)
items += [item]
Expand All @@ -136,7 +136,7 @@ async def process_mm_data_async(
audio_offsets = None
item = MultimodalDataItem(
feature=[res["audio_features"]],
audio_feature_lens=res["audio_feature_lens"],
model_specific_data={"audio_feature_lens": res["audio_feature_lens"]},
offsets=audio_offsets,
modality=Modality.AUDIO,
)
Expand Down
35 changes: 17 additions & 18 deletions python/sglang/srt/multimodal/processors/mlama.py
Original file line number Diff line number Diff line change
@@ -1,34 +1,33 @@
from typing import List, Union

from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.models.mllama import MllamaForConditionalGeneration
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.utils import load_image
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)


class MllamaImageProcessor(BaseMultimodalProcessor):
models = [MllamaForConditionalGeneration]

def __init__(self, hf_config, server_args, _processor):
super().__init__(hf_config, server_args, _processor)
self.IM_TOKEN = self._processor.image_token
self.IM_TOKEN_ID = self._processor.image_token_id

async def process_mm_data_async(
self, image_data: List[Union[str, bytes]], input_text, *args, **kwargs
):
if isinstance(input_text, list):
assert len(input_text) and isinstance(input_text[0], int)
input_text = self._processor.tokenizer.decode(input_text)
base_out = self.load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=MultimodalSpecialTokens(image_token=self.IM_TOKEN),
)

images = [load_image(image)[0] for image in image_data]
image_inputs = self.process_mm_data(input_text=input_text, images=images)
image_inputs["input_ids"] = image_inputs["input_ids"].tolist()[0]
image_inputs["mm_items"] = [
MultimodalDataItem(
feature=image_inputs["pixel_values"],
aspect_ratio_id=image_inputs["aspect_ratio_ids"],
aspect_ratio_mask=image_inputs["aspect_ratio_mask"],
modality=Modality.IMAGE,
)
]
mm_items, input_ids, _ = self.process_and_combine_mm_data(base_out)
Comment thread
JustinTong0323 marked this conversation as resolved.
Outdated

return image_inputs
return {
"mm_items": mm_items,
"input_ids": input_ids.tolist(),
"im_token_id": self.IM_TOKEN_ID,
}
Loading
Loading