From 9ee3bee79ee93623fadcb20e3d61a37cf5435023 Mon Sep 17 00:00:00 2001 From: QwertyJack <7554089+QwertyJack@users.noreply.github.com> Date: Sat, 5 Sep 2026 14:07:18 +0000 Subject: [PATCH 1/2] fix(frontend): backport sparse multimodal placeholder masks Signed-off-by: QwertyJack <7554089+QwertyJack@users.noreply.github.com> --- .../platform/test_mm_placeholder_mask.py | 161 ++++++++++++++++++ vllm_ascend/patch/__init__.py | 8 + .../platform/patch_mm_placeholder_mask.py | 90 ++++++++++ vllm_ascend/platform.py | 4 + 4 files changed, 263 insertions(+) create mode 100644 tests/ut/patch/platform/test_mm_placeholder_mask.py create mode 100644 vllm_ascend/patch/platform/patch_mm_placeholder_mask.py diff --git a/tests/ut/patch/platform/test_mm_placeholder_mask.py b/tests/ut/patch/platform/test_mm_placeholder_mask.py new file mode 100644 index 000000000000..24a0888582b9 --- /dev/null +++ b/tests/ut/patch/platform/test_mm_placeholder_mask.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch +from vllm.entrypoints.scale_out.derender.serving import ServingDerender +from vllm.entrypoints.scale_out.render.serving import ServingRender +from vllm.entrypoints.scale_out.token_in_token_out.protocol import GenerateRequest, MultiModalFeatures +from vllm.entrypoints.scale_out.token_in_token_out.serving import ServingTokens +from vllm.inputs import mm_input +from vllm.multimodal.inputs import MultiModalKwargsItems, PlaceholderRange + +from vllm_ascend.patch.platform.patch_mm_placeholder_mask import ( + _wrap_serve_tokens, + install_mm_placeholder_mask_patch, +) + + +@pytest.mark.parametrize("serving", [ServingRender, ServingDerender]) +@pytest.mark.parametrize("mask", [None, [True, False, True, True]]) +def test_render_json_preserves_sparse_and_dense_placeholders(serving, mask): + install_mm_placeholder_mask_patch() + engine_input = mm_input( + prompt_token_ids=[1, 2, 3, 4, 5], + mm_kwargs=MultiModalKwargsItems({}), + mm_hashes={"image": ["image-0"]}, + mm_placeholders={ + "image": [PlaceholderRange(offset=1, length=4, is_embed=None if mask is None else torch.tensor(mask))] + }, + ) + features = serving._extract_mm_features(engine_input) + request = GenerateRequest.model_validate_json( + GenerateRequest(token_ids=[1, 2, 3, 4, 5], sampling_params={}, features=features).model_dump_json() + ) + assert request.features.mm_placeholders["image"][0].is_embed == mask + assert serving._extract_mm_features({"type": "token", "prompt_token_ids": [1]}) is None + + +def test_installation_is_idempotent(): + install_mm_placeholder_mask_patch() + render = ServingRender._extract_mm_features + serve = ServingTokens.serve_tokens + install_mm_placeholder_mask_patch() + assert ServingRender._extract_mm_features is render + assert ServingTokens.serve_tokens is serve + assert "is_embed" in GenerateRequest.model_json_schema()["$defs"]["PlaceholderRangeInfo"]["properties"] + + +def test_fastapi_request_schema_preserves_embedding_mask(): + from fastapi import FastAPI + from fastapi.testclient import TestClient + + install_mm_placeholder_mask_patch() + app = FastAPI() + + @app.post("/generate") + def generate(request: GenerateRequest): + return request.features + + with TestClient(app) as client: + response = client.post( + "/generate", + json={ + "token_ids": [1, 2, 3, 4], + "sampling_params": {}, + "features": { + "mm_hashes": {"image": ["test"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": [True, False, True, True]}]}, + }, + }, + ) + assert response.status_code == 200 + assert response.json()["mm_placeholders"]["image"][0]["is_embed"] == [True, False, True, True] + schema = client.get("/openapi.json").json() + assert "is_embed" in schema["components"]["schemas"]["PlaceholderRangeInfo"]["properties"] + + +async def _upstream_handler(self, request, raw_request=None): + await asyncio.sleep(0) + return mm_input( + prompt_token_ids=request.token_ids, + mm_kwargs=MultiModalKwargsItems({}), + mm_hashes=request.features.mm_hashes, + mm_placeholders={"image": [PlaceholderRange(offset=0, length=4)]}, + ) + + +def test_concurrent_requests_rebuild_their_own_embedding_masks(): + install_mm_placeholder_mask_patch() + original_builder = _upstream_handler.__globals__["mm_input"] + wrapped = _wrap_serve_tokens(_upstream_handler) + + async def run(): + requests = [ + SimpleNamespace( + token_ids=[1, 2, 3, 4], + features=MultiModalFeatures.model_validate( + { + "mm_hashes": {"image": [str(i)]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": mask}]}, + } + ), + ) + for i, mask in enumerate(([True, False, True, True], [False, True, False, False], None)) + ] + return await asyncio.gather(*(wrapped(None, request) for request in requests)) + + results = asyncio.run(run()) + assert [r["mm_placeholders"]["image"][0].get_num_embeds() for r in results] == [3, 1, 4] + assert results[0]["mm_placeholders"]["image"][0].is_embed.dtype == torch.bool + assert _upstream_handler.__globals__["mm_input"] is original_builder + + +@pytest.mark.parametrize("content_parts", [None, [{"type": "image_url", "url": "image.png"}]]) +def test_text_and_content_parts_use_the_original_handler(content_parts): + async def original(self, request, raw_request): + return request + + request = SimpleNamespace(features=None if content_parts is None else object(), content_parts=content_parts) + assert asyncio.run(_wrap_serve_tokens(original)(None, request)) is request + + +def test_real_handler_passes_mask_to_engine_input_builder(): + install_mm_placeholder_mask_patch() + request = GenerateRequest.model_validate( + { + "token_ids": [1, 2, 3, 4], + "sampling_params": {}, + "features": { + "mm_hashes": {"image": ["test"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": [True, False, True, True]}]}, + }, + } + ) + + async def check_model(request): + return None + + serving = SimpleNamespace( + _check_model=check_model, + engine_client=SimpleNamespace( + errored=False, vllm_config=SimpleNamespace(scheduler_config=SimpleNamespace(max_num_seqs=4)) + ), + _maybe_get_adapters=lambda *a, **k: None, + models=SimpleNamespace(model_name=lambda _: "test"), + _base_request_id=lambda *a: "test", + ) + + class ReachedInputBuilder(Exception): + pass + + def capture(**kwargs): + assert kwargs["mm_placeholders"]["image"][0].get_num_embeds() == 3 + raise ReachedInputBuilder + + original = getattr(ServingTokens.serve_tokens, "__wrapped__", ServingTokens.serve_tokens) + with patch.dict(original.__globals__, mm_input=capture), pytest.raises(ReachedInputBuilder): + asyncio.run(ServingTokens.serve_tokens(serving, request)) diff --git a/vllm_ascend/patch/__init__.py b/vllm_ascend/patch/__init__.py index 780f2da6f4e7..ec2cc37a5e3b 100644 --- a/vllm_ascend/patch/__init__.py +++ b/vllm_ascend/patch/__init__.py @@ -27,6 +27,14 @@ # ---------------------------------------------------------------------------------- # What's Patched and how it works: +# +# platform/patch_mm_placeholder_mask.py backports +# https://github.com/vllm-project/vllm/pull/54548. It preserves sparse +# PlaceholderRange.is_embed masks through render/derender serialization and +# token-input reconstruction. Installation runs after platform/config imports, +# rebuilds nested request schemas, and skips versions with native is_embed. +# Remove this compatibility patch once every supported vLLM version preserves +# the mask. The request-local input builder leaves concurrent requests isolated. # -------------------------------- # * Platform Patch: # ================= diff --git a/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py b/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py new file mode 100644 index 000000000000..6c315798244c --- /dev/null +++ b/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py @@ -0,0 +1,90 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Compatibility backport of vllm-project/vllm#54548.""" + +from functools import wraps +from types import FunctionType + + +def _wrap_extract_features(original, placeholder_type): + @wraps(original) + def extract(engine_input): + features = original(engine_input) + if features is None: + return None + features.mm_placeholders = { + modality: [ + placeholder_type( + 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 engine_input["mm_placeholders"].items() + } + return features + + return extract + + +def _wrap_serve_tokens(original): + @wraps(original) + async def serve(self, request, raw_request=None): + if request.features is None or getattr(request, "content_parts", None): + return await original(self, request, raw_request) + + import torch + from vllm.multimodal.inputs import PlaceholderRange + + original_mm_input = original.__globals__["mm_input"] + + def mm_input(**kwargs): + kwargs["mm_placeholders"] = { + 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 request.features.mm_placeholders.items() + } + return original_mm_input(**kwargs) + + # Bind only this request's input builder without mutating module globals + # across concurrent awaits or copying the version-specific async handler. + scoped = FunctionType( + original.__code__, + dict(original.__globals__, mm_input=mm_input), + original.__name__, + original.__defaults__, + original.__closure__, + ) + scoped.__kwdefaults__ = original.__kwdefaults__ + return await scoped(self, request, raw_request) + + return serve + + +def install_mm_placeholder_mask_patch() -> None: + # Run after platform registration, not during config/plugin module imports. + from pydantic.fields import FieldInfo + from vllm.entrypoints.scale_out.derender.serving import ServingDerender + from vllm.entrypoints.scale_out.render.serving import ServingRender + from vllm.entrypoints.scale_out.token_in_token_out import protocol + from vllm.entrypoints.scale_out.token_in_token_out.serving import ServingTokens + + placeholder_type = protocol.PlaceholderRangeInfo + if "is_embed" in placeholder_type.model_fields: + return + + # Preserve class identities already imported by serving modules and routers. + placeholder_type.model_fields["is_embed"] = FieldInfo(annotation=list[bool] | None, default=None) + for model in (placeholder_type, protocol.MultiModalFeatures, protocol.GenerateRequest): + model.model_rebuild(force=True) + for serving in (ServingRender, ServingDerender): + serving._extract_mm_features = staticmethod( + _wrap_extract_features(serving._extract_mm_features, placeholder_type) + ) + ServingTokens.serve_tokens = _wrap_serve_tokens(ServingTokens.serve_tokens) diff --git a/vllm_ascend/platform.py b/vllm_ascend/platform.py index 3a80610b1511..d2f202e702ea 100644 --- a/vllm_ascend/platform.py +++ b/vllm_ascend/platform.py @@ -311,6 +311,10 @@ def pre_register_and_update(cls, parser: FlexibleArgumentParser | None = None) - register_deepseek_v4_vision_config_convertor() + from vllm_ascend.patch.platform.patch_mm_placeholder_mask import install_mm_placeholder_mask_patch + + install_mm_placeholder_mask_patch() + # For online serving, "ascend" quantization method is not a choice natively, # so we need to add "ascend" quantization method to quantization methods list # and the user can enable quantization using "vllm serve --quantization ascend". From 06ab7d806e74b6c20c3e2245b66a8e8bbd6e3104 Mon Sep 17 00:00:00 2001 From: QwertyJack <7554089+QwertyJack@users.noreply.github.com> Date: Sat, 5 Sep 2026 14:23:03 +0000 Subject: [PATCH 2/2] fix(frontend): align placeholder backport with upstream serving implementation Signed-off-by: QwertyJack <7554089+QwertyJack@users.noreply.github.com> --- .../platform/test_mm_placeholder_mask.py | 224 +++++++++++++----- vllm_ascend/patch/__init__.py | 3 +- .../platform/patch_mm_placeholder_mask.py | 194 ++++++++++++--- 3 files changed, 316 insertions(+), 105 deletions(-) diff --git a/tests/ut/patch/platform/test_mm_placeholder_mask.py b/tests/ut/patch/platform/test_mm_placeholder_mask.py index 24a0888582b9..b471522b5aec 100644 --- a/tests/ut/patch/platform/test_mm_placeholder_mask.py +++ b/tests/ut/patch/platform/test_mm_placeholder_mask.py @@ -2,20 +2,22 @@ import asyncio from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import AsyncMock, Mock, patch import pytest import torch from vllm.entrypoints.scale_out.derender.serving import ServingDerender from vllm.entrypoints.scale_out.render.serving import ServingRender +from vllm.entrypoints.scale_out.token_in_token_out import serving as upstream_serving from vllm.entrypoints.scale_out.token_in_token_out.protocol import GenerateRequest, MultiModalFeatures from vllm.entrypoints.scale_out.token_in_token_out.serving import ServingTokens from vllm.inputs import mm_input from vllm.multimodal.inputs import MultiModalKwargsItems, PlaceholderRange from vllm_ascend.patch.platform.patch_mm_placeholder_mask import ( - _wrap_serve_tokens, install_mm_placeholder_mask_patch, + rebuild_mm_placeholders, + serve_tokens, ) @@ -78,84 +80,176 @@ def generate(request: GenerateRequest): assert "is_embed" in schema["components"]["schemas"]["PlaceholderRangeInfo"]["properties"] -async def _upstream_handler(self, request, raw_request=None): - await asyncio.sleep(0) - return mm_input( - prompt_token_ids=request.token_ids, - mm_kwargs=MultiModalKwargsItems({}), - mm_hashes=request.features.mm_hashes, - mm_placeholders={"image": [PlaceholderRange(offset=0, length=4)]}, - ) - - -def test_concurrent_requests_rebuild_their_own_embedding_masks(): - install_mm_placeholder_mask_patch() - original_builder = _upstream_handler.__globals__["mm_input"] - wrapped = _wrap_serve_tokens(_upstream_handler) - - async def run(): - requests = [ - SimpleNamespace( - token_ids=[1, 2, 3, 4], - features=MultiModalFeatures.model_validate( - { - "mm_hashes": {"image": [str(i)]}, - "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": mask}]}, - } - ), - ) - for i, mask in enumerate(([True, False, True, True], [False, True, False, False], None)) - ] - return await asyncio.gather(*(wrapped(None, request) for request in requests)) - - results = asyncio.run(run()) - assert [r["mm_placeholders"]["image"][0].get_num_embeds() for r in results] == [3, 1, 4] - assert results[0]["mm_placeholders"]["image"][0].is_embed.dtype == torch.bool - assert _upstream_handler.__globals__["mm_input"] is original_builder - - -@pytest.mark.parametrize("content_parts", [None, [{"type": "image_url", "url": "image.png"}]]) -def test_text_and_content_parts_use_the_original_handler(content_parts): - async def original(self, request, raw_request): - return request - - request = SimpleNamespace(features=None if content_parts is None else object(), content_parts=content_parts) - assert asyncio.run(_wrap_serve_tokens(original)(None, request)) is request - - -def test_real_handler_passes_mask_to_engine_input_builder(): +def make_request(mask=None, **kwargs): install_mm_placeholder_mask_patch() - request = GenerateRequest.model_validate( + return GenerateRequest.model_validate( { "token_ids": [1, 2, 3, 4], - "sampling_params": {}, + "sampling_params": {"max_tokens": 8}, "features": { "mm_hashes": {"image": ["test"]}, - "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": [True, False, True, True]}]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 4, "is_embed": mask}]}, }, + **kwargs, } ) - async def check_model(request): - return None - serving = SimpleNamespace( - _check_model=check_model, +def make_serving(): + async def full_response(request, result, *args): + await asyncio.sleep(0) + return result + + return SimpleNamespace( + _check_model=AsyncMock(return_value=None), engine_client=SimpleNamespace( - errored=False, vllm_config=SimpleNamespace(scheduler_config=SimpleNamespace(max_num_seqs=4)) + errored=False, + vllm_config=SimpleNamespace(scheduler_config=SimpleNamespace(max_num_seqs=4)), + generate=Mock(side_effect=lambda engine_input, *args, **kwargs: engine_input), ), - _maybe_get_adapters=lambda *a, **k: None, + _maybe_get_adapters=Mock(return_value=None), models=SimpleNamespace(model_name=lambda _: "test"), _base_request_id=lambda *a: "test", + force_no_detokenize=False, + _log_inputs=Mock(), + _get_data_parallel_rank=Mock(return_value=2), + _get_session_id_from_headers=Mock(return_value="session-test"), + serve_tokens_full_generator=full_response, + serve_tokens_stream_generator=Mock(side_effect=lambda request, result, *args: result), + create_error_response=lambda error: error, + online_renderer=SimpleNamespace(preprocess_completion=AsyncMock(return_value=({"type": "token"},))), ) - class ReachedInputBuilder(Exception): - pass - def capture(**kwargs): - assert kwargs["mm_placeholders"]["image"][0].get_num_embeds() == 3 - raise ReachedInputBuilder +@pytest.mark.parametrize("mask", [None, [True, False, True, True]]) +def test_rebuild_mm_placeholders_restores_is_embed(mask): + request = make_request(mask) + features = MultiModalFeatures.model_validate_json(request.features.model_dump_json()) + (placeholder,) = rebuild_mm_placeholders(features.mm_placeholders)["image"] + assert placeholder.offset == 0 + assert placeholder.length == 4 + if mask is None: + assert placeholder.is_embed is None + else: + assert torch.equal(placeholder.is_embed, torch.tensor(mask, dtype=torch.bool)) + assert placeholder.is_embed.device.type == "cpu" + assert placeholder.get_num_embeds() == (4 if mask is None else sum(mask)) - original = getattr(ServingTokens.serve_tokens, "__wrapped__", ServingTokens.serve_tokens) - with patch.dict(original.__globals__, mm_input=capture), pytest.raises(ReachedInputBuilder): - asyncio.run(ServingTokens.serve_tokens(serving, request)) + +def test_concurrent_requests_rebuild_their_own_embedding_masks(): + original_builder = upstream_serving.mm_input + serving = make_serving() + + async def run(): + requests = [make_request(mask) for mask in ([True, False, True, True], [False, True, False, False], None)] + return await asyncio.gather(*(ServingTokens.serve_tokens(serving, request) for request in requests)) + + results = asyncio.run(run()) + assert [r["mm_placeholders"]["image"][0].get_num_embeds() for r in results] == [3, 1, 4] + assert results[0]["mm_placeholders"]["image"][0].is_embed.dtype == torch.bool + assert upstream_serving.mm_input is original_builder + + +def test_text_uses_original_preprocessing(): + serving = make_serving() + request = make_request(features=None) + assert asyncio.run(serve_tokens(serving, request)) == {"type": "token"} + serving.online_renderer.preprocess_completion.assert_awaited_once_with( + request, prompt_input=request.token_ids, prompt_embeds=None, skip_mm_cache=True + ) + + +@pytest.mark.skipif("content_parts" not in GenerateRequest.model_fields, reason="v0.27.1 has no content_parts") +def test_content_parts_uses_original_rendering(): + serving = make_serving() + serving.model_config = object() + serving.online_renderer.renderer = SimpleNamespace( + render_cmpl_async=AsyncMock(return_value=({"type": "multimodal"},)) + ) + request = make_request(features=None, content_parts=[{"type": "image_url", "url": "image.png"}]) + tracker = Mock() + tracker.resolve_items = AsyncMock(return_value=({"image": ["pixels"]}, {"image": ["image-id"]})) + with patch("vllm.entrypoints.chat_utils.AsyncMultiModalItemTracker", return_value=tracker): + assert asyncio.run(serve_tokens(serving, request)) == {"type": "multimodal"} + tracker.create_parser.return_value.parse_image.assert_called_once_with("image.png", None) + serving.online_renderer.renderer.render_cmpl_async.assert_awaited_once_with( + [ + { + "prompt_token_ids": request.token_ids, + "multi_modal_data": {"image": ["pixels"]}, + "multi_modal_uuids": {"image": ["image-id"]}, + } + ] + ) + + +@pytest.mark.parametrize("stream", [False, True]) +def test_engine_handoff_preserves_routing_and_output_kind(stream): + from vllm.sampling_params import RequestOutputKind + + serving = make_serving() + request = make_request([True, False, True, True], stream=stream) + old_output_kind = request.sampling_params.output_kind + result = asyncio.run(ServingTokens.serve_tokens(serving, request)) + assert result["mm_placeholders"]["image"][0].get_num_embeds() == 3 + assert result["mm_kwargs"]["image"] == [None] + kwargs = serving.engine_client.generate.call_args.kwargs + assert kwargs["data_parallel_rank"] == 2 + if "content_parts" in GenerateRequest.model_fields: + assert kwargs["session_id"] == "session-test" + assert request.sampling_params.output_kind == ( + RequestOutputKind.DELTA if stream else RequestOutputKind.FINAL_ONLY + ) + else: + assert "session_id" not in kwargs + assert request.sampling_params.output_kind == (RequestOutputKind.DELTA if stream else old_output_kind) + + +def test_model_error_does_not_schedule_request(): + serving = make_serving() + serving._check_model.return_value = "model error" + assert asyncio.run(serve_tokens(serving, make_request())) == "model error" + serving.engine_client.generate.assert_not_called() + + +def test_sampling_validation_does_not_schedule_request(): + serving = make_serving() + request = make_request(sampling_params={"n": 5, "max_tokens": 8}) + assert "max_num_seqs" in asyncio.run(serve_tokens(serving, request)) + serving.engine_client.generate.assert_not_called() + + +def test_default_max_tokens_and_no_detokenize_are_preserved(): + serving = make_serving() + serving.model_config = SimpleNamespace(max_model_len=128) + serving.default_sampling_params = {"max_tokens": 23} + serving.override_max_tokens = None + serving._extract_prompt_len = lambda _: 4 + serving.force_no_detokenize = True + request = make_request(sampling_params={}) + asyncio.run(serve_tokens(serving, request)) + assert request.sampling_params.max_tokens == 23 + assert request.sampling_params.detokenize is False + + +def test_serialized_embedding_data_reaches_engine(): + from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import encode_mm_kwargs_item + from vllm.multimodal.inputs import MultiModalBatchedField, MultiModalFieldElem, MultiModalKwargsItem + + serving = make_serving() + request = make_request([True, False, True, True]) + pixels = torch.arange(12).reshape(3, 4) + item = MultiModalKwargsItem({"pixel_values": MultiModalFieldElem(data=pixels, field=MultiModalBatchedField())}) + request.features.kwargs_data = {"image": [encode_mm_kwargs_item(item)]} + result = asyncio.run(serve_tokens(serving, request)) + assert torch.equal(result["mm_kwargs"]["image"][0]["pixel_values"].data, pixels) + assert result["mm_placeholders"]["image"][0].get_num_embeds() == 3 + + +def test_generation_exception_does_not_affect_next_request(): + serving = make_serving() + serving.engine_client.generate.side_effect = RuntimeError("generation failed") + with pytest.raises(RuntimeError, match="generation failed"): + asyncio.run(serve_tokens(serving, make_request([True, False, True, True]))) + result = asyncio.run(serve_tokens(make_serving(), make_request([False, True, False, False]))) + assert result["mm_placeholders"]["image"][0].get_num_embeds() == 1 diff --git a/vllm_ascend/patch/__init__.py b/vllm_ascend/patch/__init__.py index ec2cc37a5e3b..4afe8b7ecda6 100644 --- a/vllm_ascend/patch/__init__.py +++ b/vllm_ascend/patch/__init__.py @@ -34,7 +34,8 @@ # token-input reconstruction. Installation runs after platform/config imports, # rebuilds nested request schemas, and skips versions with native is_embed. # Remove this compatibility patch once every supported vLLM version preserves -# the mask. The request-local input builder leaves concurrent requests isolated. +# the mask. The placeholder builder and serving method follow the upstream fix; +# no function cloning or request-scoped module-global overrides are used. # -------------------------------- # * Platform Patch: # ================= diff --git a/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py b/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py index 6c315798244c..c20c138b165d 100644 --- a/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py +++ b/vllm_ascend/patch/platform/patch_mm_placeholder_mask.py @@ -1,8 +1,19 @@ # SPDX-License-Identifier: Apache-2.0 -"""Compatibility backport of vllm-project/vllm#54548.""" +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Compatibility backport of vllm-project/vllm#54548. +The placeholder builder follows upstream commit 60fe831acefe. The serving method +backports its call site onto the supported vLLM main pin ba07e4a48, with v0.27.1 +content-parts, output-kind and session-ID differences gated explicitly. +""" + +from collections.abc import Mapping from functools import wraps -from types import FunctionType +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from vllm.entrypoints.scale_out.token_in_token_out.protocol import PlaceholderRangeInfo + from vllm.multimodal.inputs import PlaceholderRange def _wrap_extract_features(original, placeholder_type): @@ -27,44 +38,149 @@ def extract(engine_input): return extract -def _wrap_serve_tokens(original): - @wraps(original) - async def serve(self, request, raw_request=None): - if request.features is None or getattr(request, "content_parts", None): - return await original(self, request, raw_request) - - import torch - from vllm.multimodal.inputs import PlaceholderRange - - original_mm_input = original.__globals__["mm_input"] - - def mm_input(**kwargs): - kwargs["mm_placeholders"] = { - 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 request.features.mm_placeholders.items() - } - return original_mm_input(**kwargs) - - # Bind only this request's input builder without mutating module globals - # across concurrent awaits or copying the version-specific async handler. - scoped = FunctionType( - original.__code__, - dict(original.__globals__, mm_input=mm_input), - original.__name__, - original.__defaults__, - original.__closure__, +def rebuild_mm_placeholders( + mm_placeholders: Mapping[str, list["PlaceholderRangeInfo"]], +) -> dict[str, list["PlaceholderRange"]]: + """Convert rendered placeholders back to ranges, as in vLLM #54548.""" + import torch + from vllm.multimodal.inputs import PlaceholderRange + + 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() + } + + +async def serve_tokens(self, request, raw_request=None): + # Keep imports lazy: this patch module can load during platform discovery, + # before vLLM config and serving modules have finished initializing. + import msgspec + from vllm.entrypoints.chat_utils import AsyncMultiModalItemTracker + from vllm.entrypoints.openai.engine.protocol import RequestResponseMetadata + from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import decode_mm_kwargs_item + from vllm.entrypoints.scale_out.token_in_token_out.serving import logger + from vllm.entrypoints.serve.utils.api_utils import get_max_tokens + from vllm.inputs import TokensPrompt, mm_input + from vllm.multimodal.inputs import MultiModalKwargsItems + from vllm.sampling_params import RequestOutputKind + + from vllm_ascend.utils import vllm_version_is + + is_v0271 = vllm_version_is("0.27.1") + error_check_ret = await self._check_model(request) + if error_check_ret is not None: + logger.error("Error with model %s", error_check_ret) + return error_check_ret + + # Preserve upstream validation before constructing or scheduling inputs. + if self.engine_client.errored: + raise self.engine_client.dead_error + + lora_request = self._maybe_get_adapters(request, supports_default_mm_loras=True) + model_name = self.models.model_name(lora_request) + request_id = f"generate-tokens-{self._base_request_id(raw_request, request.request_id)}" + request_metadata = RequestResponseMetadata(request_id=request_id) + if raw_request: + raw_request.state.request_metadata = request_metadata + + sampling_params = request.sampling_params + max_num_seqs = self.engine_client.vllm_config.scheduler_config.max_num_seqs + if sampling_params.n > max_num_seqs: + return self.create_error_response( + f"sampling_params.n must be at most the server's max_num_seqs ({max_num_seqs}), got {sampling_params.n}." + ) + try: + msgspec.msgpack.encode(sampling_params) + except (OverflowError, TypeError, ValueError) as e: + return self.create_error_response(e) + + if not is_v0271 and request.content_parts: + tracker = AsyncMultiModalItemTracker(self.model_config) + mm_parser = tracker.create_parser() + for part in request.content_parts: + ptype = part.get("type", "") + url = part.get("url") + uuid = part.get("uuid") + if ptype == "image_url": + mm_parser.parse_image(url, uuid) + elif ptype == "audio_url": + mm_parser.parse_audio(url, uuid) + elif ptype == "video_url": + mm_parser.parse_video(url, uuid) + mm_data, mm_uuids = await tracker.resolve_items() + prompt = TokensPrompt(prompt_token_ids=request.token_ids) + if mm_data: + prompt["multi_modal_data"] = mm_data + if mm_uuids: + prompt["multi_modal_uuids"] = mm_uuids + (engine_input,) = await self.online_renderer.renderer.render_cmpl_async([prompt]) + elif features := request.features: + mm_placeholders = rebuild_mm_placeholders(features.mm_placeholders) + + # Deserialize tensor data when present; None means an encoder-cache hit. + mm_kwargs = {} + if features.kwargs_data is not None: + for modality, items in features.kwargs_data.items(): + mm_kwargs[modality] = [decode_mm_kwargs_item(item) if item is not None else None for item in items] + else: + for modality, hashes in features.mm_hashes.items(): + mm_kwargs[modality] = [None] * len(hashes) + + engine_input = mm_input( + prompt_token_ids=request.token_ids, + mm_kwargs=MultiModalKwargsItems(mm_kwargs), + mm_hashes=features.mm_hashes, + mm_placeholders=mm_placeholders, + cache_salt=request.cache_salt, ) - scoped.__kwdefaults__ = original.__kwdefaults__ - return await scoped(self, request, raw_request) + else: + (engine_input,) = await self.online_renderer.preprocess_completion( + request, prompt_input=request.token_ids, prompt_embeds=None, skip_mm_cache=True + ) + + # Retain upstream defaults, logging, routing and response generators. + if not request.is_sampling_param_provided("max_tokens"): + sampling_params.max_tokens = get_max_tokens( + max_model_len=self.model_config.max_model_len, + max_tokens=None, + input_length=self._extract_prompt_len(engine_input), + default_sampling_params=self.default_sampling_params, + override_max_tokens=self.override_max_tokens, + ) + + if self.force_no_detokenize: + sampling_params.detokenize = False + if request.stream: + sampling_params.output_kind = RequestOutputKind.DELTA + elif not is_v0271: + sampling_params.output_kind = RequestOutputKind.FINAL_ONLY - return serve + self._log_inputs(request_id, engine_input, params=sampling_params, lora_request=lora_request) + trace_headers = None if raw_request is None else await self._get_trace_headers(raw_request.headers) + data_parallel_rank = self._get_data_parallel_rank(raw_request) + # Session routing is not part of the v0.27.1 engine-client interface. + session_kwargs = {} if is_v0271 else {"session_id": self._get_session_id_from_headers(raw_request)} + result_generator = self.engine_client.generate( + engine_input, + sampling_params, + request_id, + lora_request=lora_request, + trace_headers=trace_headers, + priority=request.priority, + data_parallel_rank=data_parallel_rank, + **session_kwargs, + ) + assert result_generator is not None + if request.stream: + return self.serve_tokens_stream_generator(request, result_generator, request_id, model_name, request_metadata) + return await self.serve_tokens_full_generator(request, result_generator, request_id, model_name, request_metadata) def install_mm_placeholder_mask_patch() -> None: @@ -87,4 +203,4 @@ def install_mm_placeholder_mask_patch() -> None: serving._extract_mm_features = staticmethod( _wrap_extract_features(serving._extract_mm_features, placeholder_type) ) - ServingTokens.serve_tokens = _wrap_serve_tokens(ServingTokens.serve_tokens) + ServingTokens.serve_tokens = serve_tokens