From 48d0a137d5d5e15e2a45a7e70ccfb1ec77e65daa Mon Sep 17 00:00:00 2001 From: roytman Date: Mon, 25 May 2026 16:10:20 +0300 Subject: [PATCH 1/2] add image_grid_thw as an input parameter to Decode and Prefill in EPD deployments Signed-off-by: roytman --- .../disagg/test_serving_multimodal_tokens.py | 10 ++++++++++ vllm/entrypoints/serve/disagg/protocol.py | 9 +++++++++ vllm/entrypoints/serve/disagg/serving.py | 20 +++++++++++++++++++ vllm/entrypoints/serve/render/serving.py | 20 ++++++++++++++++++- 4 files changed, 58 insertions(+), 1 deletion(-) diff --git a/tests/entrypoints/serve/disagg/test_serving_multimodal_tokens.py b/tests/entrypoints/serve/disagg/test_serving_multimodal_tokens.py index e13dd425af14..cd619cf677cb 100644 --- a/tests/entrypoints/serve/disagg/test_serving_multimodal_tokens.py +++ b/tests/entrypoints/serve/disagg/test_serving_multimodal_tokens.py @@ -123,6 +123,16 @@ async def test_render_to_generate_roundtrip(client, test_image): assert "image" in features["kwargs_data"] assert len(features["kwargs_data"]["image"]) > 0 + assert "image_grid_thw" in features + assert "image" in features["image_grid_thw"] + image_grids = features["image_grid_thw"]["image"] + assert isinstance(image_grids, list) + assert len(image_grids) == len(features["kwargs_data"]["image"]) + for grid in image_grids: + assert isinstance(grid, list) + assert len(grid) == 3 + assert all(isinstance(v, int) and v > 0 for v in grid) + # Build generate request from render output generate_payload = render_data generate_payload["sampling_params"] = { diff --git a/vllm/entrypoints/serve/disagg/protocol.py b/vllm/entrypoints/serve/disagg/protocol.py index 60d2a6424a00..f393a81b05c9 100644 --- a/vllm/entrypoints/serve/disagg/protocol.py +++ b/vllm/entrypoints/serve/disagg/protocol.py @@ -58,6 +58,15 @@ class MultiModalFeatures(BaseModel): ``None`` for metadata-only (cache-hit) responses. """ + image_grid_thw: dict[str, list[list[int] | None]] | None = None + """Per-modality grid dimensions for mRoPE position computation. + + Each value is a list parallel to ``mm_hashes[modality]``. Each entry + is a ``[grid_t, grid_h, grid_w]`` list (or ``None`` for cache hits). + This allows the prefill worker to compute mRoPE positions without + deserializing the full ``kwargs_data`` blobs. + """ + class GenerateRequest(BaseModel): request_id: str = Field( diff --git a/vllm/entrypoints/serve/disagg/serving.py b/vllm/entrypoints/serve/disagg/serving.py index 0cc227ee74db..95b8a1542676 100644 --- a/vllm/entrypoints/serve/disagg/serving.py +++ b/vllm/entrypoints/serve/disagg/serving.py @@ -11,6 +11,7 @@ import msgspec import numpy as np import pybase64 as base64 +import torch from fastapi import Request from vllm.engine.protocol import EngineClient @@ -43,6 +44,8 @@ from vllm.logger import init_logger from vllm.logprobs import Logprob from vllm.multimodal.inputs import ( + MultiModalBatchedField, + MultiModalFieldElem, MultiModalKwargsItem, MultiModalKwargsItems, PlaceholderRange, @@ -156,6 +159,23 @@ async def serve_tokens( decode_mm_kwargs_item(item) if item is not None else None for item in items ] + elif features.image_grid_thw is not None: + # Lightweight path: construct minimal items containing + # only grid metadata for mRoPE position computation. + for modality, grids in features.image_grid_thw.items(): + items_list: list[MultiModalKwargsItem | None] = [] + thw_key = f"{modality}_grid_thw" + for grid in grids: + if grid is not None: + tensor = torch.tensor([grid], dtype=torch.int64) + elem = MultiModalFieldElem( + data=tensor, + field=MultiModalBatchedField(keep_on_cpu=True), + ) + items_list.append(MultiModalKwargsItem({thw_key: elem})) + else: + items_list.append(None) + mm_kwargs[modality] = items_list else: for modality, hashes in features.mm_hashes.items(): mm_kwargs[modality] = [None] * len(hashes) diff --git a/vllm/entrypoints/serve/render/serving.py b/vllm/entrypoints/serve/render/serving.py index 782b2eaea24b..1cc693cb0be2 100644 --- a/vllm/entrypoints/serve/render/serving.py +++ b/vllm/entrypoints/serve/render/serving.py @@ -2,10 +2,13 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Sequence from http import HTTPStatus -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast from openai_harmony import Message as OpenAIMessage +if TYPE_CHECKING: + import torch + from vllm.config import ModelConfig from vllm.entrypoints.chat_utils import ( ChatTemplateContentFormatOption, @@ -380,18 +383,33 @@ def _extract_mm_features( # Serialize tensor data per modality. kwargs_data: dict[str, list[str | None]] | None = None + image_grid_thw: dict[str, list[list[int] | None]] | None = None if raw_mm_kwargs := mm_engine_input.get("mm_kwargs"): kwargs_data = {} + image_grid_thw = {} for modality, items in raw_mm_kwargs.items(): kwargs_data[modality] = [ encode_mm_kwargs_item(item) if item is not None else None for item in items ] + thw_key = f"{modality}_grid_thw" + grids: list[list[int] | None] = [] + for item in items: + if item is not None and thw_key in item: + thw_tensor = cast("torch.Tensor", item[thw_key].data) + grids.append(thw_tensor.tolist()) + else: + grids.append(None) + if any(g is not None for g in grids): + image_grid_thw[modality] = grids + if not image_grid_thw: + image_grid_thw = None return MultiModalFeatures( mm_hashes=mm_hashes, mm_placeholders=mm_placeholders, kwargs_data=kwargs_data, + image_grid_thw=image_grid_thw, ) def _make_request_with_harmony( From 2b75ae62054135a0d6c9e4cf06396a25bb944d56 Mon Sep 17 00:00:00 2001 From: roytman Date: Mon, 25 May 2026 20:19:04 +0300 Subject: [PATCH 2/2] fix comments Signed-off-by: roytman --- vllm/entrypoints/serve/disagg/serving.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vllm/entrypoints/serve/disagg/serving.py b/vllm/entrypoints/serve/disagg/serving.py index 95b8a1542676..3ffa652f39ff 100644 --- a/vllm/entrypoints/serve/disagg/serving.py +++ b/vllm/entrypoints/serve/disagg/serving.py @@ -167,7 +167,7 @@ async def serve_tokens( thw_key = f"{modality}_grid_thw" for grid in grids: if grid is not None: - tensor = torch.tensor([grid], dtype=torch.int64) + tensor = torch.tensor(grid, dtype=torch.int64) elem = MultiModalFieldElem( data=tensor, field=MultiModalBatchedField(keep_on_cpu=True),