Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -212,8 +212,8 @@ def test_placeholder_ranges_from_engine_input():

assert placeholders is not None
assert [item.model_dump() for item in placeholders["image"]] == [
{"offset": 1, "length": 5},
{"offset": 8, "length": 2},
{"offset": 1, "length": 5, "is_embed": None},
{"offset": 8, "length": 2, "is_embed": None},
]

text_only: Any = {"type": "token", "prompt_token_ids": [1, 2, 3]}
Expand Down Expand Up @@ -586,7 +586,7 @@ async def mock_generate(*args, **kwargs):
"prompt_token_ids": [1, 2, 2, 2, 3],
"mm_placeholders": {"image": [PlaceholderRange(offset=1, length=3)]},
}
MM_PLACEHOLDERS = {"image": [{"offset": 1, "length": 3}]}
MM_PLACEHOLDERS = {"image": [{"offset": 1, "length": 3, "is_embed": None}]}


def _build_mm_serving_tokens(engine: AsyncLLM) -> ServingTokens:
Expand Down
90 changes: 90 additions & 0 deletions tests/entrypoints/scale_out/token_in_token_out/test_mm_serde.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
MultiModalFieldElem,
MultiModalFlatField,
MultiModalKwargsItem,
MultiModalKwargsItems,
MultiModalSharedField,
PlaceholderRange,
)
Expand Down Expand Up @@ -370,3 +371,92 @@ def test_metadata_can_replace_or_extend_full_mm_data():
merged = merge_mm_kwargs_items(full_item, metadata_item)
assert merged is not None
assert set(merged) == {"pixel_values", "image_grid_thw"}


def test_render_features_preserve_is_embed():
"""A sparse `is_embed` mask must reach the generate side intact.

The model runner branches on `PlaceholderRange.is_embed`: with a mask a
span consumes only `is_embed.sum()` rows of encoder output and only those
positions are marked as embeddings; without one it consumes the whole span
and every position in it is overwritten. Dropping the mask on the wire
therefore corrupts the prompt for models that use sparse placeholders
(Gemma 3, Phi-3-V, the Qwen omni thinkers).
"""
is_embed = torch.tensor([True, False, True, True], dtype=torch.bool)
engine_input = mm_input(
prompt_token_ids=[1, 2, 3, 4, 5],
mm_kwargs=MultiModalKwargsItems({}),
mm_hashes={"image": ["hash-0"]},
mm_placeholders={
"image": [PlaceholderRange(offset=1, length=4, is_embed=is_embed)]
},
)

features = extract_mm_features(engine_input)
assert features is not None

# Assert on the serialized form: this is what actually crosses the wire
# to the generate service.
dumped = features.model_dump()["mm_placeholders"]["image"][0]
assert dumped.get("is_embed") == [True, False, True, True]


def test_rebuild_mm_placeholders_restores_is_embed():
from vllm.entrypoints.scale_out.token_in_token_out.mm_features import (
rebuild_mm_placeholders,
)

is_embed = [True, False, True, True]
features = MultiModalFeatures(
mm_hashes={"image": ["hash-0"]},
mm_placeholders={
"image": [PlaceholderRangeInfo(offset=1, length=4, is_embed=is_embed)]
},
)
features = MultiModalFeatures.model_validate(features.model_dump())

(placeholder,) = rebuild_mm_placeholders(features.mm_placeholders)["image"]
assert placeholder.offset == 1
assert placeholder.length == 4
assert placeholder.is_embed is not None
assert torch.equal(placeholder.is_embed, torch.tensor(is_embed, dtype=torch.bool))
# The row count the model runner slices the encoder output by.
assert placeholder.get_num_embeds() == 3


def test_placeholder_without_is_embed_roundtrips_as_none():
from vllm.entrypoints.scale_out.token_in_token_out.mm_features import (
rebuild_mm_placeholders,
)

engine_input = mm_input(
prompt_token_ids=[1, 2, 3],
mm_kwargs=MultiModalKwargsItems({}),
mm_hashes={"image": ["hash-0"]},
mm_placeholders={"image": [PlaceholderRange(offset=0, length=3)]},
)

features = extract_mm_features(engine_input)
assert features is not None
assert features.mm_placeholders["image"][0].is_embed is None

(placeholder,) = rebuild_mm_placeholders(features.mm_placeholders)["image"]
assert placeholder.is_embed is None
assert placeholder.get_num_embeds() == 3


def test_placeholder_rejects_a_mask_shorter_than_the_span():
"""The wire format has to check what ``PlaceholderRange`` does not.

``is_embed`` is documented as a mask of shape ``(length,)`` and the
frozen dataclass never validates it, so a client-supplied short mask
would reach ``get_embeds_indices_in_range`` and slice the wrong rows
out of the encoder output.
"""
try:
PlaceholderRangeInfo(offset=0, length=4, is_embed=[True])
except ValidationError as exc:
assert "is_embed has 1 entries" in str(exc)
return
raise AssertionError("expected a mask shorter than the span to fail")
37 changes: 35 additions & 2 deletions vllm/entrypoints/scale_out/token_in_token_out/mm_features.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,11 @@

from __future__ import annotations

from collections.abc import Callable, Collection, Sequence
from collections.abc import Callable, Collection, Mapping, Sequence
from typing import cast

import torch

from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import (
decode_mm_kwargs_item,
encode_mm_kwargs_item,
Expand All @@ -24,6 +26,7 @@
from vllm.multimodal.inputs import (
MultiModalKwargsItem,
MultiModalKwargsOptionalItems,
PlaceholderRange,
)


Expand Down Expand Up @@ -117,6 +120,31 @@ def mm_kwargs_from_features(
return mm_kwargs


def rebuild_mm_placeholders(
mm_placeholders: Mapping[str, list[PlaceholderRangeInfo]],
) -> dict[str, list[PlaceholderRange]]:
"""Convert rendered `PlaceholderRangeInfo` back into `PlaceholderRange`.

`is_embed` has to survive the round trip: the model runner branches on it
to decide how much of the encoder output a span consumes and which
positions to overwrite, and it cannot be recomputed from offset and
length alone.
"""
return {
modality: [
PlaceholderRange(
offset=p.offset,
length=p.length,
is_embed=None
if p.is_embed is None
else torch.tensor(p.is_embed, dtype=torch.bool),
)
for p in ranges
]
for modality, ranges in mm_placeholders.items()
}


def placeholder_ranges_from_engine_input(
engine_input: EngineInput,
) -> dict[str, list[PlaceholderRangeInfo]] | None:
Expand All @@ -129,7 +157,12 @@ def placeholder_ranges_from_engine_input(
]
return {
modality: [
PlaceholderRangeInfo(offset=p.offset, length=p.length) for p in ranges
PlaceholderRangeInfo(
offset=p.offset,
length=p.length,
is_embed=None if p.is_embed is None else p.is_embed.tolist(),
)
for p in ranges
]
for modality, ranges in raw_placeholders.items()
}
Expand Down
19 changes: 16 additions & 3 deletions vllm/entrypoints/scale_out/token_in_token_out/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,9 +48,22 @@ class PlaceholderRangeInfo(BaseModel):
length: int = Field(gt=0)
"""Number of placeholder tokens."""

# TODO: add `is_embed: list[bool] | None` once the /generate side
# consumes features — some models (e.g. Qwen-VL) use sparse
# placeholder masks that cannot be recomputed from offset+length alone.
is_embed: list[bool] | None = None
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"""Which positions in the span actually receive embeddings.

``None`` means every position does. Models with sparse placeholder
masks (Gemma 3, Phi-3-V, the Qwen omni thinkers) mark only a subset,
and the mask cannot be recomputed from offset and length alone.
"""

@model_validator(mode="after")
def _check_is_embed_length(self) -> "PlaceholderRangeInfo":
if self.is_embed is not None and len(self.is_embed) != self.length:
raise ValueError(
f"is_embed has {len(self.is_embed)} entries but the "
f"placeholder spans {self.length} tokens"
)
return self


def _has_serialized_mm_items(
Expand Down
10 changes: 2 additions & 8 deletions vllm/entrypoints/scale_out/token_in_token_out/serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,6 @@
from vllm.multimodal.inputs import (
MultiModalKwargsItem,
MultiModalKwargsItems,
PlaceholderRange,
)
from vllm.outputs import RequestOutput
from vllm.renderers.online_renderer import OnlineRenderer
Expand All @@ -55,6 +54,7 @@
from .mm_features import (
mm_kwargs_from_features,
placeholder_ranges_from_engine_input,
rebuild_mm_placeholders,
)
from .protocol import (
GenerateRequest,
Expand Down Expand Up @@ -254,13 +254,7 @@ async def serve_tokens(
[prompt]
)
elif features := request.features:
# Convert PlaceholderRangeInfo → PlaceholderRange per modality.
mm_placeholders: dict[str, list[PlaceholderRange]] = {
modality: [
PlaceholderRange(offset=p.offset, length=p.length) for p in ranges
]
for modality, ranges in features.mm_placeholders.items()
}
mm_placeholders = rebuild_mm_placeholders(features.mm_placeholders)

# Deserialize full tensor data and optional metadata-only data.
# Metadata-only items are valid when ec_transfer_params is set.
Expand Down
Loading