[VLM] Enable per-image MM splitting by default and remove MULTI_IMAGES modality (#21899)
This commit is contained in:
@@ -467,9 +467,6 @@ class Envs:
|
|||||||
SGLANG_MM_FEATURE_CACHE_MB = EnvInt(4 * 1024)
|
SGLANG_MM_FEATURE_CACHE_MB = EnvInt(4 * 1024)
|
||||||
SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05)
|
SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05)
|
||||||
|
|
||||||
# MM splitting behavior control
|
|
||||||
SGLANG_ENABLE_MM_SPLITTING = EnvBool(False)
|
|
||||||
|
|
||||||
# Mamba
|
# Mamba
|
||||||
SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16")
|
SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16")
|
||||||
SGLANG_MAMBA_SSM_DTYPE = EnvStr(None)
|
SGLANG_MAMBA_SSM_DTYPE = EnvStr(None)
|
||||||
|
|||||||
@@ -327,15 +327,13 @@ class MultiModalityDataPaddingPatternMultimodalTokens(MultiModalityDataPaddingPa
|
|||||||
|
|
||||||
input_ids_tensor = torch.as_tensor(input_ids)
|
input_ids_tensor = torch.as_tensor(input_ids)
|
||||||
|
|
||||||
# Check if MM splitting is enabled
|
# Replace multimodal tokens using per-item offsets
|
||||||
if envs.SGLANG_ENABLE_MM_SPLITTING.get():
|
|
||||||
items_by_modality = defaultdict(list)
|
items_by_modality = defaultdict(list)
|
||||||
for item in mm_inputs.mm_items:
|
for item in mm_inputs.mm_items:
|
||||||
items_by_modality[item.modality].append(item)
|
items_by_modality[item.modality].append(item)
|
||||||
|
|
||||||
token_id_map = {
|
token_id_map = {
|
||||||
Modality.IMAGE: mm_inputs.im_token_id,
|
Modality.IMAGE: mm_inputs.im_token_id,
|
||||||
Modality.MULTI_IMAGES: mm_inputs.im_token_id,
|
|
||||||
Modality.AUDIO: mm_inputs.audio_token_id,
|
Modality.AUDIO: mm_inputs.audio_token_id,
|
||||||
Modality.VIDEO: mm_inputs.video_token_id,
|
Modality.VIDEO: mm_inputs.video_token_id,
|
||||||
}
|
}
|
||||||
@@ -349,24 +347,6 @@ class MultiModalityDataPaddingPatternMultimodalTokens(MultiModalityDataPaddingPa
|
|||||||
for i, item in enumerate(items):
|
for i, item in enumerate(items):
|
||||||
for offset in items[i].offsets:
|
for offset in items[i].offsets:
|
||||||
input_ids_tensor[offset[0] : offset[1] + 1] = item.pad_value
|
input_ids_tensor[offset[0] : offset[1] + 1] = item.pad_value
|
||||||
else:
|
|
||||||
# Create mapping of token_ids to pad_values for each modality
|
|
||||||
token_to_pad_mapping = {}
|
|
||||||
for item in mm_inputs.mm_items:
|
|
||||||
if item.is_image() and mm_inputs.im_token_id is not None:
|
|
||||||
token_to_pad_mapping[mm_inputs.im_token_id] = item.pad_value
|
|
||||||
elif item.is_audio() and mm_inputs.audio_token_id is not None:
|
|
||||||
token_to_pad_mapping[mm_inputs.audio_token_id] = item.pad_value
|
|
||||||
elif item.is_video() and mm_inputs.video_token_id is not None:
|
|
||||||
token_to_pad_mapping[mm_inputs.video_token_id] = item.pad_value
|
|
||||||
else:
|
|
||||||
raise ValueError(
|
|
||||||
f"No multimodal token id provided for {item.modality}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Apply replacements for all tokens at once
|
|
||||||
for token_id, pad_value in token_to_pad_mapping.items():
|
|
||||||
input_ids_tensor[input_ids_tensor == token_id] = pad_value
|
|
||||||
|
|
||||||
ret_input_ids = input_ids_tensor.tolist()
|
ret_input_ids = input_ids_tensor.tolist()
|
||||||
return ret_input_ids
|
return ret_input_ids
|
||||||
@@ -1476,6 +1456,54 @@ def _slice_model_data(
|
|||||||
return sliced
|
return sliced
|
||||||
|
|
||||||
|
|
||||||
|
def _try_simple_split(item, num_items, expanded_mm_items):
|
||||||
|
"""Try to split a bundled item by matching feature dim-0 to offset count.
|
||||||
|
Returns True if split succeeded, False otherwise."""
|
||||||
|
feature = item.feature if item.feature is not None else item.precomputed_embeddings
|
||||||
|
if feature is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if isinstance(feature, (torch.Tensor, np.ndarray)):
|
||||||
|
feature_count = feature.shape[0]
|
||||||
|
elif isinstance(feature, (list, tuple)):
|
||||||
|
feature_count = len(feature)
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if feature_count != num_items:
|
||||||
|
return False
|
||||||
|
|
||||||
|
for i in range(num_items):
|
||||||
|
new_item = copy.copy(item)
|
||||||
|
if item.feature is not None:
|
||||||
|
if isinstance(item.feature, (list, tuple)):
|
||||||
|
new_item.feature = [item.feature[i]]
|
||||||
|
else:
|
||||||
|
new_item.feature = item.feature[i : i + 1]
|
||||||
|
if item.precomputed_embeddings is not None:
|
||||||
|
if isinstance(item.precomputed_embeddings, (list, tuple)):
|
||||||
|
new_item.precomputed_embeddings = [item.precomputed_embeddings[i]]
|
||||||
|
else:
|
||||||
|
new_item.precomputed_embeddings = item.precomputed_embeddings[i : i + 1]
|
||||||
|
new_item.offsets = [item.offsets[i]]
|
||||||
|
new_data = {}
|
||||||
|
for k, v in item.model_specific_data.items():
|
||||||
|
if isinstance(v, (list, tuple)) and len(v) == num_items:
|
||||||
|
new_data[k] = [v[i]]
|
||||||
|
elif (
|
||||||
|
isinstance(v, (torch.Tensor, np.ndarray))
|
||||||
|
and len(v.shape) > 0
|
||||||
|
and v.shape[0] == num_items
|
||||||
|
):
|
||||||
|
new_data[k] = v[i : i + 1]
|
||||||
|
else:
|
||||||
|
new_data[k] = v
|
||||||
|
new_item.model_specific_data = new_data
|
||||||
|
new_item.hash = None
|
||||||
|
expanded_mm_items.append(new_item)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def get_new_expanded_mm_items(original_mm_items):
|
def get_new_expanded_mm_items(original_mm_items):
|
||||||
expanded_mm_items = []
|
expanded_mm_items = []
|
||||||
for item in original_mm_items:
|
for item in original_mm_items:
|
||||||
@@ -1488,6 +1516,8 @@ def get_new_expanded_mm_items(original_mm_items):
|
|||||||
image_grid_thw = item.model_specific_data.get("image_grid_thw")
|
image_grid_thw = item.model_specific_data.get("image_grid_thw")
|
||||||
grid_len = _get_length(image_grid_thw)
|
grid_len = _get_length(image_grid_thw)
|
||||||
if image_grid_thw is None or grid_len != num_items:
|
if image_grid_thw is None or grid_len != num_items:
|
||||||
|
# No grid info — fall back to simple split by feature dim-0
|
||||||
|
if not _try_simple_split(item, num_items, expanded_mm_items):
|
||||||
expanded_mm_items.append(item)
|
expanded_mm_items.append(item)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -1533,6 +1563,7 @@ def get_new_expanded_mm_items(original_mm_items):
|
|||||||
elif item.is_video():
|
elif item.is_video():
|
||||||
video_grid_thw = item.model_specific_data.get("video_grid_thw")
|
video_grid_thw = item.model_specific_data.get("video_grid_thw")
|
||||||
if video_grid_thw is None:
|
if video_grid_thw is None:
|
||||||
|
if not _try_simple_split(item, num_items, expanded_mm_items):
|
||||||
expanded_mm_items.append(item)
|
expanded_mm_items.append(item)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -1623,6 +1654,7 @@ def get_new_expanded_mm_items(original_mm_items):
|
|||||||
new_item.hash = None
|
new_item.hash = None
|
||||||
expanded_mm_items.append(new_item)
|
expanded_mm_items.append(new_item)
|
||||||
else:
|
else:
|
||||||
|
if not _try_simple_split(item, num_items, expanded_mm_items):
|
||||||
expanded_mm_items.append(item)
|
expanded_mm_items.append(item)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -198,7 +198,6 @@ class FINISH_ABORT(BaseFinishReason):
|
|||||||
|
|
||||||
class Modality(Enum):
|
class Modality(Enum):
|
||||||
IMAGE = auto()
|
IMAGE = auto()
|
||||||
MULTI_IMAGES = auto()
|
|
||||||
VIDEO = auto()
|
VIDEO = auto()
|
||||||
AUDIO = auto()
|
AUDIO = auto()
|
||||||
|
|
||||||
@@ -225,9 +224,10 @@ class MultimodalInputFormat(Enum):
|
|||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class MultimodalDataItem:
|
class MultimodalDataItem:
|
||||||
"""
|
"""
|
||||||
One MultimodalDataItem contains all inputs for one modality.
|
One MultimodalDataItem represents a single multimodal input (one image, one video, or one audio).
|
||||||
For example, if there are 3 images and 1 audio inputs, there will be 2 MultimodalDataItem.
|
For example, if there are 3 images and 1 audio, there will be 4 MultimodalDataItems.
|
||||||
One for images and one for audio.
|
|
||||||
|
Each item has its own hash and pad_value, enabling per-image RadixAttention caching.
|
||||||
|
|
||||||
We put the common fields first and the model-specific fields in model_specific_data.
|
We put the common fields first and the model-specific fields in model_specific_data.
|
||||||
"""
|
"""
|
||||||
@@ -305,7 +305,7 @@ class MultimodalDataItem:
|
|||||||
return self.modality == Modality.AUDIO
|
return self.modality == Modality.AUDIO
|
||||||
|
|
||||||
def is_image(self):
|
def is_image(self):
|
||||||
return self.modality in [Modality.IMAGE, Modality.MULTI_IMAGES]
|
return self.modality == Modality.IMAGE
|
||||||
|
|
||||||
def is_video(self):
|
def is_video(self):
|
||||||
return self.modality == Modality.VIDEO
|
return self.modality == Modality.VIDEO
|
||||||
@@ -330,12 +330,6 @@ class MultimodalDataItem:
|
|||||||
ret.validate()
|
ret.validate()
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def merge(self, other):
|
|
||||||
self.feature += other.feature
|
|
||||||
self.offsets += other.offsets
|
|
||||||
self.hash = hash((self.hash, other.hash))
|
|
||||||
self.set_pad_value()
|
|
||||||
|
|
||||||
def reconstruct(self):
|
def reconstruct(self):
|
||||||
if not isinstance(self.feature, CudaIpcTensorTransportProxy):
|
if not isinstance(self.feature, CudaIpcTensorTransportProxy):
|
||||||
return
|
return
|
||||||
@@ -395,19 +389,10 @@ class MultimodalInputs:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dict(obj: dict):
|
def from_dict(obj: dict):
|
||||||
original_mm_items = obj["mm_items"]
|
mm_items = obj["mm_items"]
|
||||||
for mm_item in original_mm_items:
|
for mm_item in mm_items:
|
||||||
mm_item.reconstruct()
|
mm_item.reconstruct()
|
||||||
|
|
||||||
# Check if MM splitting is enabled
|
|
||||||
if not envs.SGLANG_ENABLE_MM_SPLITTING.get():
|
|
||||||
mm_items = original_mm_items
|
|
||||||
else:
|
|
||||||
from sglang.srt.managers.mm_utils import get_new_expanded_mm_items
|
|
||||||
|
|
||||||
# Now, `mm_items` contains one item per image.
|
|
||||||
mm_items = get_new_expanded_mm_items(original_mm_items)
|
|
||||||
|
|
||||||
ret = MultimodalInputs(
|
ret = MultimodalInputs(
|
||||||
mm_items=mm_items,
|
mm_items=mm_items,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -58,6 +58,21 @@ _KNOWN_BROKEN_AUTOMODEL_ERROR = "Could not find VoxtralRealtimeTextModel"
|
|||||||
|
|
||||||
|
|
||||||
class LlavaBaseForCausalLM(nn.Module):
|
class LlavaBaseForCausalLM(nn.Module):
|
||||||
|
@staticmethod
|
||||||
|
def _infer_image_aspect_ratio(mm_items):
|
||||||
|
"""Determine image_aspect_ratio from processor metadata or item count."""
|
||||||
|
# Check if processor stored the aspect_ratio it used
|
||||||
|
for item in mm_items:
|
||||||
|
ar = item.model_specific_data.get("image_aspect_ratio")
|
||||||
|
if ar is not None:
|
||||||
|
return ar
|
||||||
|
# Fallback: multi-image or video → pad, single image → anyres
|
||||||
|
image_items = [item for item in mm_items if item.is_image()]
|
||||||
|
has_video = any(item.is_video() for item in mm_items)
|
||||||
|
if len(image_items) > 1 or has_video:
|
||||||
|
return "pad"
|
||||||
|
return "anyres"
|
||||||
|
|
||||||
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
|
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
|
||||||
image_sizes = flatten_nested_list(
|
image_sizes = flatten_nested_list(
|
||||||
[item.image_sizes for item in image_inputs.mm_items]
|
[item.image_sizes for item in image_inputs.mm_items]
|
||||||
@@ -66,13 +81,8 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
pad_values = [item.pad_value for item in image_inputs.mm_items]
|
pad_values = [item.pad_value for item in image_inputs.mm_items]
|
||||||
|
|
||||||
# hardcode for spatial_unpad + anyres
|
# hardcode for spatial_unpad + anyres
|
||||||
if any(
|
# Use per-item aspect_ratio from processor if available, else infer
|
||||||
item.modality == Modality.MULTI_IMAGES or item.modality == Modality.VIDEO
|
image_aspect_ratio = self._infer_image_aspect_ratio(image_inputs.mm_items)
|
||||||
for item in image_inputs.mm_items
|
|
||||||
):
|
|
||||||
image_aspect_ratio = "pad"
|
|
||||||
else:
|
|
||||||
image_aspect_ratio = "anyres"
|
|
||||||
offset_list = []
|
offset_list = []
|
||||||
image_inputs.image_pad_len = []
|
image_inputs.image_pad_len = []
|
||||||
for image_idx, image_s in enumerate(image_sizes):
|
for image_idx, image_s in enumerate(image_sizes):
|
||||||
@@ -168,13 +178,9 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
# Embed text inputs
|
# Embed text inputs
|
||||||
input_embeds = self.language_model.model.embed_tokens(input_ids)
|
input_embeds = self.language_model.model.embed_tokens(input_ids)
|
||||||
|
|
||||||
# Got List[List[str]] extend it to List[str]
|
# Compute max image offset per request to determine need_vision
|
||||||
# The length of the List should be equal to batch size
|
|
||||||
modalities_list = []
|
|
||||||
max_image_offset = []
|
max_image_offset = []
|
||||||
for im in image_inputs:
|
for im in image_inputs:
|
||||||
if im:
|
|
||||||
modalities_list.extend([item.modality for item in im.mm_items])
|
|
||||||
if im and im.image_offsets:
|
if im and im.image_offsets:
|
||||||
max_image_offset.append(
|
max_image_offset.append(
|
||||||
np.max(np.array(im.image_offsets) + np.array(im.image_pad_len))
|
np.max(np.array(im.image_offsets) + np.array(im.image_pad_len))
|
||||||
@@ -187,6 +193,18 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
|
|
||||||
if need_vision.any():
|
if need_vision.any():
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|
||||||
|
# Build per-image lists filtered by need_vision
|
||||||
|
modalities_list = []
|
||||||
|
aspect_ratios = [] # per-image aspect ratio
|
||||||
|
for i in range(bs):
|
||||||
|
if need_vision[i] and image_inputs[i]:
|
||||||
|
items = image_inputs[i].mm_items
|
||||||
|
ar = self._infer_image_aspect_ratio(items)
|
||||||
|
for item in items:
|
||||||
|
modalities_list.append(item.modality)
|
||||||
|
aspect_ratios.append(ar)
|
||||||
|
|
||||||
pixel_values = flatten_nested_list(
|
pixel_values = flatten_nested_list(
|
||||||
[
|
[
|
||||||
[item.feature for item in image_inputs[i].mm_items]
|
[item.feature for item in image_inputs[i].mm_items]
|
||||||
@@ -194,12 +212,12 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
if need_vision[i]
|
if need_vision[i]
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
# Per-image sizes (each entry is [(w,h)] for one image)
|
||||||
image_sizes = [
|
image_sizes = [
|
||||||
flatten_nested_list(
|
item.image_sizes
|
||||||
[item.image_sizes for item in image_inputs[i].mm_items]
|
|
||||||
)
|
|
||||||
for i in range(bs)
|
for i in range(bs)
|
||||||
if need_vision[i]
|
if need_vision[i]
|
||||||
|
for item in image_inputs[i].mm_items
|
||||||
]
|
]
|
||||||
|
|
||||||
########## Encode Image ########
|
########## Encode Image ########
|
||||||
@@ -228,18 +246,7 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
new_image_features = []
|
new_image_features = []
|
||||||
height = width = self.num_patches_per_side
|
height = width = self.num_patches_per_side
|
||||||
for image_idx, image_feature in enumerate(image_features):
|
for image_idx, image_feature in enumerate(image_features):
|
||||||
if modalities_list[image_idx] == Modality.IMAGE:
|
image_aspect_ratio = aspect_ratios[image_idx]
|
||||||
image_aspect_ratio = (
|
|
||||||
self.config.image_aspect_ratio
|
|
||||||
) # single image
|
|
||||||
elif (
|
|
||||||
modalities_list[image_idx] == Modality.MULTI_IMAGES
|
|
||||||
or modalities_list[image_idx] == Modality.VIDEO
|
|
||||||
):
|
|
||||||
image_aspect_ratio = "pad" # multi image
|
|
||||||
# image_aspect_ratio = (
|
|
||||||
# "anyres" if len(image_sizes[image_idx]) == 1 else "pad"
|
|
||||||
# )
|
|
||||||
if (
|
if (
|
||||||
image_feature.shape[0] > 1
|
image_feature.shape[0] > 1
|
||||||
and "anyres" in image_aspect_ratio
|
and "anyres" in image_aspect_ratio
|
||||||
@@ -388,6 +395,7 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
extend_start_loc_cpu = forward_batch.extend_start_loc.cpu().numpy()
|
extend_start_loc_cpu = forward_batch.extend_start_loc.cpu().numpy()
|
||||||
extend_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
|
extend_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
|
||||||
prefix_lens_cpu = forward_batch.extend_prefix_lens_cpu
|
prefix_lens_cpu = forward_batch.extend_prefix_lens_cpu
|
||||||
|
# Fill in the image features using flat indexing (one pt per image)
|
||||||
pt = 0
|
pt = 0
|
||||||
for i in range(bs):
|
for i in range(bs):
|
||||||
if not need_vision[i]:
|
if not need_vision[i]:
|
||||||
@@ -396,20 +404,25 @@ class LlavaBaseForCausalLM(nn.Module):
|
|||||||
start_idx = extend_start_loc_cpu[i]
|
start_idx = extend_start_loc_cpu[i]
|
||||||
seq_len = extend_seq_lens[i]
|
seq_len = extend_seq_lens[i]
|
||||||
prefix_len = prefix_lens_cpu[i]
|
prefix_len = prefix_lens_cpu[i]
|
||||||
|
n_images = len(image_inputs[i].image_offsets)
|
||||||
|
|
||||||
|
for j in range(n_images):
|
||||||
|
image_offset = image_inputs[i].image_offsets[j]
|
||||||
|
|
||||||
# Multiple images
|
|
||||||
for image_idx, image_offset in enumerate(
|
|
||||||
image_inputs[i].image_offsets
|
|
||||||
):
|
|
||||||
if (
|
if (
|
||||||
image_offset + image_inputs[i].image_pad_len[image_idx]
|
image_offset + image_inputs[i].image_pad_len[j]
|
||||||
<= prefix_len
|
<= prefix_len
|
||||||
):
|
):
|
||||||
|
pt += 1
|
||||||
continue
|
continue
|
||||||
if image_offset >= prefix_len + seq_len:
|
if image_offset >= prefix_len + seq_len:
|
||||||
|
pt += n_images - j
|
||||||
break
|
break
|
||||||
|
|
||||||
tmp_image_feature = image_features[pt][image_idx]
|
tmp_image_feature = image_features[pt]
|
||||||
|
# Squeeze batch dim from per-image features [1, feat, hidden]
|
||||||
|
if tmp_image_feature.ndim == 3:
|
||||||
|
tmp_image_feature = tmp_image_feature[0]
|
||||||
pad_len = tmp_image_feature.shape[0]
|
pad_len = tmp_image_feature.shape[0]
|
||||||
|
|
||||||
input_offset = image_offset - prefix_len
|
input_offset = image_offset - prefix_len
|
||||||
|
|||||||
@@ -993,7 +993,11 @@ class MiniCPMV2_6(MiniCPMBaseModel):
|
|||||||
slice_end_id: int = image_inputs.slice_end_id
|
slice_end_id: int = image_inputs.slice_end_id
|
||||||
|
|
||||||
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
||||||
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
|
# Only increment data_idx on im_start (not slice_start) so all slices
|
||||||
|
# within one image share the same pad_value for per-image caching.
|
||||||
|
pattern = MultiModalityDataPaddingPatternTokenPairs(
|
||||||
|
media_token_pairs, data_start_token_ids=[im_start_id]
|
||||||
|
)
|
||||||
|
|
||||||
return pattern.pad_input_tokens(input_ids, image_inputs)
|
return pattern.pad_input_tokens(input_ids, image_inputs)
|
||||||
|
|
||||||
@@ -1155,7 +1159,11 @@ class MiniCPMV4_0(MiniCPMBaseModel):
|
|||||||
slice_end_id: int = image_inputs.slice_end_id
|
slice_end_id: int = image_inputs.slice_end_id
|
||||||
|
|
||||||
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
||||||
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
|
# Only increment data_idx on im_start (not slice_start) so all slices
|
||||||
|
# within one image share the same pad_value for per-image caching.
|
||||||
|
pattern = MultiModalityDataPaddingPatternTokenPairs(
|
||||||
|
media_token_pairs, data_start_token_ids=[im_start_id]
|
||||||
|
)
|
||||||
|
|
||||||
return pattern.pad_input_tokens(input_ids, image_inputs)
|
return pattern.pad_input_tokens(input_ids, image_inputs)
|
||||||
|
|
||||||
@@ -1321,7 +1329,11 @@ class MiniCPMV4_5(MiniCPMBaseModel):
|
|||||||
slice_end_id: int = image_inputs.slice_end_id
|
slice_end_id: int = image_inputs.slice_end_id
|
||||||
|
|
||||||
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
|
||||||
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
|
# Only increment data_idx on im_start (not slice_start) so all slices
|
||||||
|
# within one image share the same pad_value for per-image caching.
|
||||||
|
pattern = MultiModalityDataPaddingPatternTokenPairs(
|
||||||
|
media_token_pairs, data_start_token_ids=[im_start_id]
|
||||||
|
)
|
||||||
|
|
||||||
return pattern.pad_input_tokens(input_ids, image_inputs)
|
return pattern.pad_input_tokens(input_ids, image_inputs)
|
||||||
|
|
||||||
|
|||||||
@@ -137,7 +137,6 @@ class MultimodalSpecialTokens:
|
|||||||
def get_token_id_by_modality(self, modality: Modality) -> Optional[int]:
|
def get_token_id_by_modality(self, modality: Modality) -> Optional[int]:
|
||||||
return {
|
return {
|
||||||
Modality.IMAGE: self.image_token_id,
|
Modality.IMAGE: self.image_token_id,
|
||||||
Modality.MULTI_IMAGES: self.image_token_id,
|
|
||||||
Modality.VIDEO: self.video_token_id,
|
Modality.VIDEO: self.video_token_id,
|
||||||
Modality.AUDIO: self.audio_token_id,
|
Modality.AUDIO: self.audio_token_id,
|
||||||
}.get(modality)
|
}.get(modality)
|
||||||
@@ -359,7 +358,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
mm_items.append(
|
mm_items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
modality=modality,
|
modality=modality,
|
||||||
offsets=offset,
|
offsets=[offset],
|
||||||
precomputed_embeddings=embedding_slice,
|
precomputed_embeddings=embedding_slice,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -998,7 +997,8 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self, data_dict: dict, modality: Modality = None
|
self, data_dict: dict, modality: Modality = None
|
||||||
) -> List[MultimodalDataItem]:
|
) -> List[MultimodalDataItem]:
|
||||||
"""
|
"""
|
||||||
Create mm_items directly from processor output, with one item for each modality
|
Create mm_items from processor output. Initially creates one item per modality;
|
||||||
|
these are later split into per-image/video items by get_new_expanded_mm_items.
|
||||||
|
|
||||||
Note that the data_dict can be passed via offline engine api
|
Note that the data_dict can be passed via offline engine api
|
||||||
"""
|
"""
|
||||||
@@ -1141,6 +1141,11 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
mm_token_id=mm_token_id,
|
mm_token_id=mm_token_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Split bundled items into per-image/video items for better cache granularity
|
||||||
|
from sglang.srt.managers.mm_utils import get_new_expanded_mm_items
|
||||||
|
|
||||||
|
all_collected_items = get_new_expanded_mm_items(all_collected_items)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
solution for cuda-ipc memory-leak:
|
solution for cuda-ipc memory-leak:
|
||||||
1. memory-pool: each time get a slice from memory-pool and use it as transport-data (with async lock guard)
|
1. memory-pool: each time get a slice from memory-pool and use it as transport-data (with async lock guard)
|
||||||
|
|||||||
@@ -588,11 +588,21 @@ class InternVLProcessor(BaseMultimodalProcessor):
|
|||||||
|
|
||||||
items = []
|
items = []
|
||||||
if image_tensor is not None:
|
if image_tensor is not None:
|
||||||
|
# Split per-image for better cache granularity
|
||||||
|
assert len(num_patches_list) == len(image_offsets), (
|
||||||
|
f"InternVL: num_patches_list ({len(num_patches_list)}) != "
|
||||||
|
f"image_offsets ({len(image_offsets)})"
|
||||||
|
)
|
||||||
|
cumulative = 0
|
||||||
|
for i, num_patches in enumerate(num_patches_list):
|
||||||
items.append(
|
items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
feature=image_tensor, modality=Modality.IMAGE, offsets=image_offsets
|
feature=image_tensor[cumulative : cumulative + num_patches],
|
||||||
|
modality=Modality.IMAGE,
|
||||||
|
offsets=[image_offsets[i]],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
cumulative += num_patches
|
||||||
if video_tensor is not None:
|
if video_tensor is not None:
|
||||||
items.append(
|
items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
@@ -702,11 +712,21 @@ class InternVLProcessor(BaseMultimodalProcessor):
|
|||||||
|
|
||||||
items = []
|
items = []
|
||||||
if pixel_values is not None:
|
if pixel_values is not None:
|
||||||
|
# Split per-image for better cache granularity
|
||||||
|
assert len(num_patches_list) == len(image_offsets), (
|
||||||
|
f"InternVL: num_patches_list ({len(num_patches_list)}) != "
|
||||||
|
f"image_offsets ({len(image_offsets)})"
|
||||||
|
)
|
||||||
|
cumulative = 0
|
||||||
|
for i, num_patches in enumerate(num_patches_list):
|
||||||
items.append(
|
items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
feature=pixel_values, modality=Modality.IMAGE, offsets=image_offsets
|
feature=pixel_values[cumulative : cumulative + num_patches],
|
||||||
|
modality=Modality.IMAGE,
|
||||||
|
offsets=[image_offsets[i]],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
cumulative += num_patches
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"input_ids": input_ids,
|
"input_ids": input_ids,
|
||||||
|
|||||||
@@ -187,34 +187,39 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
|
|||||||
pixel_values.append(pixel_v)
|
pixel_values.append(pixel_v)
|
||||||
data_hashes.append(image_h)
|
data_hashes.append(image_h)
|
||||||
image_sizes.append(image_s)
|
image_sizes.append(image_s)
|
||||||
|
|
||||||
if isinstance(pixel_values[0], np.ndarray):
|
|
||||||
pixel_values = np.stack(pixel_values, axis=0)
|
|
||||||
else:
|
else:
|
||||||
# A single image
|
# A single image
|
||||||
pixel_values, image_hash, image_size = await self._process_single_image(
|
pixel_values, image_hash, image_size = await self._process_single_image(
|
||||||
image_data[0], aspect_ratio, grid_pinpoints
|
image_data[0], aspect_ratio, grid_pinpoints
|
||||||
)
|
)
|
||||||
|
pixel_values = [pixel_values]
|
||||||
image_sizes = [image_size]
|
image_sizes = [image_size]
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid image data: {image_data}")
|
raise ValueError(f"Invalid image data: {image_data}")
|
||||||
modality = Modality.IMAGE
|
modality = Modality.IMAGE
|
||||||
if isinstance(request_obj.modalities, list):
|
if isinstance(request_obj.modalities, list):
|
||||||
if request_obj.modalities[0] == "multi-images":
|
if request_obj.modalities[0] == "video":
|
||||||
modality = Modality.MULTI_IMAGES
|
|
||||||
elif request_obj.modalities[0] == "video":
|
|
||||||
modality = Modality.VIDEO
|
modality = Modality.VIDEO
|
||||||
|
|
||||||
return {
|
# Create one item per image for better cache granularity
|
||||||
"mm_items": [
|
mm_items = []
|
||||||
|
for pixel_v, image_s in zip(pixel_values, image_sizes):
|
||||||
|
# Ensure ndim=4 so the model forward takes the correct encode branch
|
||||||
|
if isinstance(pixel_v, np.ndarray) and pixel_v.ndim == 3:
|
||||||
|
pixel_v = np.expand_dims(pixel_v, 0)
|
||||||
|
mm_items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
feature=pixel_values,
|
feature=pixel_v,
|
||||||
model_specific_data={
|
model_specific_data={
|
||||||
"image_sizes": image_sizes,
|
"image_sizes": [image_s],
|
||||||
|
"image_aspect_ratio": aspect_ratio,
|
||||||
},
|
},
|
||||||
modality=modality,
|
modality=modality,
|
||||||
)
|
)
|
||||||
],
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"mm_items": mm_items,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -223,6 +223,8 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
|
|||||||
f"{len(pixel_values)} vs. {len(tgt_sizes)}"
|
f"{len(pixel_values)} vs. {len(tgt_sizes)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Track slices per image (like vLLM's num_slices)
|
||||||
|
slices_per_image: List[int] = []
|
||||||
pixel_values_flat: List[torch.Tensor] = []
|
pixel_values_flat: List[torch.Tensor] = []
|
||||||
tgt_sizes_flat: List[torch.Tensor] = []
|
tgt_sizes_flat: List[torch.Tensor] = []
|
||||||
for pixel_b, tgt_b in zip(pixel_values, tgt_sizes):
|
for pixel_b, tgt_b in zip(pixel_values, tgt_sizes):
|
||||||
@@ -231,6 +233,7 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Inconsistent N lengths, found: " f"{len(pixel_b)} vs {len(tgt_b)}"
|
"Inconsistent N lengths, found: " f"{len(pixel_b)} vs {len(tgt_b)}"
|
||||||
)
|
)
|
||||||
|
slices_per_image.append(len(pixel_b))
|
||||||
for pixel_n, tgt_n in zip(pixel_b, tgt_b):
|
for pixel_n, tgt_n in zip(pixel_b, tgt_b):
|
||||||
pixel_values_flat += [pixel_n]
|
pixel_values_flat += [pixel_n]
|
||||||
tgt_sizes_flat += [tgt_n]
|
tgt_sizes_flat += [tgt_n]
|
||||||
@@ -250,14 +253,23 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
|
|||||||
image_offsets.extend(slice_offsets)
|
image_offsets.extend(slice_offsets)
|
||||||
image_offsets = sorted(image_offsets)
|
image_offsets = sorted(image_offsets)
|
||||||
|
|
||||||
|
# Create one item per image, each with its own slices and offsets
|
||||||
if len(pixel_values) != 0:
|
if len(pixel_values) != 0:
|
||||||
item = MultimodalDataItem(
|
pv_idx = 0
|
||||||
feature=pixel_values,
|
offset_idx = 0
|
||||||
offsets=image_offsets,
|
for num_slices in slices_per_image:
|
||||||
model_specific_data={"tgt_size": tgt_sizes_flat},
|
items.append(
|
||||||
|
MultimodalDataItem(
|
||||||
|
feature=pixel_values[pv_idx : pv_idx + num_slices],
|
||||||
|
offsets=image_offsets[offset_idx : offset_idx + num_slices],
|
||||||
|
model_specific_data={
|
||||||
|
"tgt_size": tgt_sizes_flat[pv_idx : pv_idx + num_slices]
|
||||||
|
},
|
||||||
modality=Modality.IMAGE,
|
modality=Modality.IMAGE,
|
||||||
)
|
)
|
||||||
items += [item]
|
)
|
||||||
|
pv_idx += num_slices
|
||||||
|
offset_idx += num_slices
|
||||||
|
|
||||||
if (
|
if (
|
||||||
"audio_features" in res
|
"audio_features" in res
|
||||||
|
|||||||
@@ -61,7 +61,7 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor):
|
|||||||
mm_items.append(
|
mm_items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
modality=modality,
|
modality=modality,
|
||||||
offsets=offset,
|
offsets=[offset],
|
||||||
precomputed_embeddings=embedding_slice,
|
precomputed_embeddings=embedding_slice,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -469,7 +469,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
mm_items.append(
|
mm_items.append(
|
||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
modality=modality,
|
modality=modality,
|
||||||
offsets=offset,
|
offsets=[offset],
|
||||||
precomputed_embeddings=embedding_slice,
|
precomputed_embeddings=embedding_slice,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -34,13 +34,13 @@ class TestMultimodalInputsFromDict(unittest.TestCase):
|
|||||||
schedule_batch.torch.cuda, "is_available", return_value=True
|
schedule_batch.torch.cuda, "is_available", return_value=True
|
||||||
), patch.object(
|
), patch.object(
|
||||||
schedule_batch.torch.cuda, "current_device", return_value=0
|
schedule_batch.torch.cuda, "current_device", return_value=0
|
||||||
), patch.object(
|
|
||||||
schedule_batch.envs.SGLANG_ENABLE_MM_SPLITTING, "get", return_value=False
|
|
||||||
), patch.object(
|
), patch.object(
|
||||||
schedule_batch.envs.SGLANG_MM_BUFFER_SIZE_MB, "get", return_value=0
|
schedule_batch.envs.SGLANG_MM_BUFFER_SIZE_MB, "get", return_value=0
|
||||||
):
|
):
|
||||||
mm_inputs = MultimodalInputs.from_dict({"mm_items": [mm_item]})
|
mm_inputs = MultimodalInputs.from_dict({"mm_items": [mm_item]})
|
||||||
|
|
||||||
|
# Splitting happens at the processor layer, not in from_dict.
|
||||||
|
# from_dict just reconstructs and passes through.
|
||||||
self.assertEqual(len(mm_inputs.mm_items), 1)
|
self.assertEqual(len(mm_inputs.mm_items), 1)
|
||||||
self.assertTrue(torch.equal(mm_inputs.mm_items[0].feature, feature_tensor))
|
self.assertTrue(torch.equal(mm_inputs.mm_items[0].feature, feature_tensor))
|
||||||
proxy_feature.reconstruct_on_target_device.assert_called_once_with(0)
|
proxy_feature.reconstruct_on_target_device.assert_called_once_with(0)
|
||||||
|
|||||||
Reference in New Issue
Block a user