diff --git a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py index 6b6d732f07e2..f14a0ca1bed1 100644 --- a/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py +++ b/tests/entrypoints/scale_out/token_in_token_out/test_generate_stream.py @@ -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]} @@ -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: diff --git a/tests/entrypoints/scale_out/token_in_token_out/test_mm_serde.py b/tests/entrypoints/scale_out/token_in_token_out/test_mm_serde.py index 5cd60a959fb3..aaf2707c5cb6 100644 --- a/tests/entrypoints/scale_out/token_in_token_out/test_mm_serde.py +++ b/tests/entrypoints/scale_out/token_in_token_out/test_mm_serde.py @@ -28,6 +28,7 @@ MultiModalFieldElem, MultiModalFlatField, MultiModalKwargsItem, + MultiModalKwargsItems, MultiModalSharedField, PlaceholderRange, ) @@ -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") diff --git a/vllm/entrypoints/scale_out/token_in_token_out/mm_features.py b/vllm/entrypoints/scale_out/token_in_token_out/mm_features.py index 62d006783a2e..017395986a86 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/mm_features.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/mm_features.py @@ -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, @@ -24,6 +26,7 @@ from vllm.multimodal.inputs import ( MultiModalKwargsItem, MultiModalKwargsOptionalItems, + PlaceholderRange, ) @@ -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: @@ -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() } diff --git a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py index 31bd7e2eaef9..8a69eec8ca46 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py @@ -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 + """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( diff --git a/vllm/entrypoints/scale_out/token_in_token_out/serving.py b/vllm/entrypoints/scale_out/token_in_token_out/serving.py index a30e519cb22e..47093f6e9f7c 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/serving.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/serving.py @@ -43,7 +43,6 @@ from vllm.multimodal.inputs import ( MultiModalKwargsItem, MultiModalKwargsItems, - PlaceholderRange, ) from vllm.outputs import RequestOutput from vllm.renderers.online_renderer import OnlineRenderer @@ -55,6 +54,7 @@ from .mm_features import ( mm_kwargs_from_features, placeholder_ranges_from_engine_input, + rebuild_mm_placeholders, ) from .protocol import ( GenerateRequest, @@ -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.