From d117a4d1a915c7acf124563f9b337de5c0aa2c2f Mon Sep 17 00:00:00 2001 From: Cyrus Leung Date: Thu, 22 Jan 2026 20:44:22 +0800 Subject: [PATCH] [Frontend] Introduce Renderer for processing chat messages (using `ModelConfig`) (#30200) Signed-off-by: DarkLight1337 --- .buildkite/test-amd.yaml | 2 + .buildkite/test-pipeline.yaml | 2 + .buildkite/test_areas/misc.yaml | 2 + docs/features/reasoning_outputs.md | 3 +- .../entrypoints/openai/test_chat_template.py | 156 ----- tests/entrypoints/openai/test_serving_chat.py | 116 ++-- .../entrypoints/openai/test_serving_engine.py | 71 --- .../entrypoints/openai/test_serving_models.py | 1 + .../openai/test_serving_responses.py | 2 + tests/entrypoints/pooling/score/test_utils.py | 8 +- tests/entrypoints/test_chat_utils.py | 492 +------------- tests/renderers/__init__.py | 0 tests/renderers/test_hf.py | 537 ++++++++++++++++ tests/renderers/test_mistral.py | 100 +++ tests/test_inputs.py | 5 +- tests/v1/engine/test_llm_engine.py | 2 +- .../engine/test_process_multi_modal_uuids.py | 21 +- tools/pre_commit/mypy.py | 1 + vllm/engine/protocol.py | 11 +- vllm/entrypoints/chat_utils.py | 535 +--------------- vllm/entrypoints/llm.py | 68 +- vllm/entrypoints/openai/api_server.py | 6 +- .../openai/chat_completion/serving.py | 16 +- vllm/entrypoints/openai/completion/serving.py | 14 +- vllm/entrypoints/openai/engine/serving.py | 163 ++--- vllm/entrypoints/openai/models/serving.py | 1 + vllm/entrypoints/openai/responses/context.py | 5 +- vllm/entrypoints/openai/responses/serving.py | 13 +- vllm/entrypoints/pooling/__init__.py | 8 +- vllm/entrypoints/pooling/classify/serving.py | 6 +- vllm/entrypoints/pooling/embed/serving.py | 6 +- vllm/entrypoints/pooling/pooling/serving.py | 9 +- vllm/entrypoints/pooling/score/serving.py | 6 +- vllm/entrypoints/pooling/score/utils.py | 9 +- vllm/entrypoints/serve/tokenize/serving.py | 12 +- vllm/entrypoints/utils.py | 43 -- vllm/inputs/preprocess.py | 15 +- vllm/renderers/__init__.py | 7 + vllm/renderers/deepseek_v32.py | 119 ++++ vllm/renderers/grok2.py | 119 ++++ vllm/renderers/hf.py | 600 ++++++++++++++++++ vllm/renderers/mistral.py | 147 +++++ vllm/renderers/protocol.py | 48 ++ vllm/renderers/registry.py | 88 +++ vllm/renderers/terratorch.py | 85 +++ vllm/v1/engine/async_llm.py | 18 +- vllm/v1/engine/input_processor.py | 12 +- vllm/v1/engine/llm_engine.py | 16 +- 48 files changed, 2141 insertions(+), 1585 deletions(-) delete mode 100644 tests/entrypoints/openai/test_chat_template.py delete mode 100644 tests/entrypoints/openai/test_serving_engine.py create mode 100644 tests/renderers/__init__.py create mode 100644 tests/renderers/test_hf.py create mode 100644 tests/renderers/test_mistral.py create mode 100644 vllm/renderers/__init__.py create mode 100644 vllm/renderers/deepseek_v32.py create mode 100644 vllm/renderers/grok2.py create mode 100644 vllm/renderers/hf.py create mode 100644 vllm/renderers/mistral.py create mode 100644 vllm/renderers/protocol.py create mode 100644 vllm/renderers/registry.py create mode 100644 vllm/renderers/terratorch.py diff --git a/.buildkite/test-amd.yaml b/.buildkite/test-amd.yaml index e87aed02731..4df8dc5f8e5 100644 --- a/.buildkite/test-amd.yaml +++ b/.buildkite/test-amd.yaml @@ -71,6 +71,7 @@ steps: - tests/test_inputs.py - tests/test_outputs.py - tests/multimodal + - tests/renderers - tests/standalone_tests/lazy_imports.py - tests/tokenizers_ - tests/tool_parsers @@ -82,6 +83,7 @@ steps: - pytest -v -s test_inputs.py - pytest -v -s test_outputs.py - pytest -v -s -m 'cpu_test' multimodal + - pytest -v -s renderers - pytest -v -s tokenizers_ - pytest -v -s tool_parsers - pytest -v -s transformers_utils diff --git a/.buildkite/test-pipeline.yaml b/.buildkite/test-pipeline.yaml index b5b919c17c7..46dfea41186 100644 --- a/.buildkite/test-pipeline.yaml +++ b/.buildkite/test-pipeline.yaml @@ -64,6 +64,7 @@ steps: - tests/test_inputs.py - tests/test_outputs.py - tests/multimodal + - tests/renderers - tests/standalone_tests/lazy_imports.py - tests/tokenizers_ - tests/tool_parsers @@ -75,6 +76,7 @@ steps: - pytest -v -s test_inputs.py - pytest -v -s test_outputs.py - pytest -v -s -m 'cpu_test' multimodal + - pytest -v -s renderers - pytest -v -s tokenizers_ - pytest -v -s tool_parsers - pytest -v -s transformers_utils diff --git a/.buildkite/test_areas/misc.yaml b/.buildkite/test_areas/misc.yaml index 252af1e56a1..b3b4566abff 100644 --- a/.buildkite/test_areas/misc.yaml +++ b/.buildkite/test_areas/misc.yaml @@ -121,6 +121,7 @@ steps: - tests/test_inputs.py - tests/test_outputs.py - tests/multimodal + - tests/renderers - tests/standalone_tests/lazy_imports.py - tests/tokenizers_ - tests/tool_parsers @@ -132,6 +133,7 @@ steps: - pytest -v -s test_inputs.py - pytest -v -s test_outputs.py - pytest -v -s -m 'cpu_test' multimodal + - pytest -v -s renderers - pytest -v -s tokenizers_ - pytest -v -s tool_parsers - pytest -v -s transformers_utils diff --git a/docs/features/reasoning_outputs.md b/docs/features/reasoning_outputs.md index 107d1d2b5bc..2bb7eeb311f 100644 --- a/docs/features/reasoning_outputs.md +++ b/docs/features/reasoning_outputs.md @@ -254,7 +254,8 @@ You can add a new `ReasoningParser` similar to [vllm/reasoning/deepseek_r1_reaso # import the required packages from vllm.reasoning import ReasoningParser, ReasoningParserManager - from vllm.entrypoints.openai.protocol import ChatCompletionRequest, DeltaMessage + from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest + from vllm.entrypoints.openai.engine.protocol import DeltaMessage # define a reasoning parser and register it to vllm # the name list in register_module can be used diff --git a/tests/entrypoints/openai/test_chat_template.py b/tests/entrypoints/openai/test_chat_template.py deleted file mode 100644 index 961ad40ca2c..00000000000 --- a/tests/entrypoints/openai/test_chat_template.py +++ /dev/null @@ -1,156 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -import pytest - -from vllm.config import ModelConfig -from vllm.entrypoints.chat_utils import apply_hf_chat_template, load_chat_template -from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest -from vllm.tokenizers import get_tokenizer - -from ...models.registry import HF_EXAMPLE_MODELS -from ...utils import VLLM_PATH - -chatml_jinja_path = VLLM_PATH / "examples/template_chatml.jinja" -assert chatml_jinja_path.exists() - -# Define models, templates, and their corresponding expected outputs -MODEL_TEMPLATE_GENERATION_OUTPUT = [ - ( - "facebook/opt-125m", - chatml_jinja_path, - True, - False, - """<|im_start|>user -Hello<|im_end|> -<|im_start|>assistant -Hi there!<|im_end|> -<|im_start|>user -What is the capital of<|im_end|> -<|im_start|>assistant -""", - ), - ( - "facebook/opt-125m", - chatml_jinja_path, - False, - False, - """<|im_start|>user -Hello<|im_end|> -<|im_start|>assistant -Hi there!<|im_end|> -<|im_start|>user -What is the capital of""", - ), - ( - "facebook/opt-125m", - chatml_jinja_path, - False, - True, - """<|im_start|>user -Hello<|im_end|> -<|im_start|>assistant -Hi there!<|im_end|> -<|im_start|>user -What is the capital of<|im_end|> -<|im_start|>assistant -The capital of""", - ), -] - -TEST_MESSAGES = [ - {"role": "user", "content": "Hello"}, - {"role": "assistant", "content": "Hi there!"}, - {"role": "user", "content": "What is the capital of"}, -] -ASSISTANT_MESSAGE_TO_CONTINUE = {"role": "assistant", "content": "The capital of"} - - -def test_load_chat_template(): - # Testing chatml template - template_content = load_chat_template(chat_template=chatml_jinja_path) - - # Test assertions - assert template_content is not None - # Hard coded value for template_chatml.jinja - assert ( - template_content - == """{% for message in messages %}{{'<|im_start|>' + message['role'] + '\\n' + message['content']}}{% if (loop.last and add_generation_prompt) or not loop.last %}{{ '<|im_end|>' + '\\n'}}{% endif %}{% endfor %} -{% if add_generation_prompt and messages[-1]['role'] != 'assistant' %}{{ '<|im_start|>assistant\\n' }}{% endif %}""" # noqa: E501 - ) - - -def test_no_load_chat_template_filelike(): - # Testing chatml template - template = "../../examples/does_not_exist" - - with pytest.raises(ValueError, match="looks like a file path"): - load_chat_template(chat_template=template) - - -def test_no_load_chat_template_literallike(): - # Testing chatml template - template = "{{ messages }}" - - template_content = load_chat_template(chat_template=template) - - assert template_content == template - - -@pytest.mark.parametrize( - "model,template,add_generation_prompt,continue_final_message,expected_output", - MODEL_TEMPLATE_GENERATION_OUTPUT, -) -def test_get_gen_prompt( - model, template, add_generation_prompt, continue_final_message, expected_output -): - model_info = HF_EXAMPLE_MODELS.find_hf_info(model) - model_info.check_available_online(on_fail="skip") - - model_config = ModelConfig( - model, - tokenizer=model_info.tokenizer or model, - tokenizer_mode=model_info.tokenizer_mode, - trust_remote_code=model_info.trust_remote_code, - revision=model_info.revision, - hf_overrides=model_info.hf_overrides, - skip_tokenizer_init=model_info.require_embed_inputs, - enable_prompt_embeds=model_info.require_embed_inputs, - enable_mm_embeds=model_info.require_embed_inputs, - enforce_eager=model_info.enforce_eager, - dtype=model_info.dtype, - ) - - # Initialize the tokenizer - tokenizer = get_tokenizer( - tokenizer_name=model_config.tokenizer, - trust_remote_code=model_config.trust_remote_code, - ) - template_content = load_chat_template(chat_template=template) - - # Create a mock request object using keyword arguments - mock_request = ChatCompletionRequest( - model=model, - messages=TEST_MESSAGES + [ASSISTANT_MESSAGE_TO_CONTINUE] - if continue_final_message - else TEST_MESSAGES, - add_generation_prompt=add_generation_prompt, - continue_final_message=continue_final_message, - ) - - # Call the function and get the result - result = apply_hf_chat_template( - tokenizer=tokenizer, - conversation=mock_request.messages, - chat_template=mock_request.chat_template or template_content, - model_config=model_config, - tools=None, - add_generation_prompt=mock_request.add_generation_prompt, - continue_final_message=mock_request.continue_final_message, - ) - - # Test assertion - assert result == expected_output, ( - f"The generated prompt does not match the expected output for " - f"model {model} and template {template}" - ) diff --git a/tests/entrypoints/openai/test_serving_chat.py b/tests/entrypoints/openai/test_serving_chat.py index dca5512c015..0f8de343557 100644 --- a/tests/entrypoints/openai/test_serving_chat.py +++ b/tests/entrypoints/openai/test_serving_chat.py @@ -11,7 +11,7 @@ import pytest_asyncio from openai import OpenAI from vllm._aiter_ops import is_aiter_found_and_supported -from vllm.config.multimodal import MultiModalConfig +from vllm.config import MultiModalConfig from vllm.entrypoints.openai.chat_completion.protocol import ( ChatCompletionRequest, ChatCompletionResponse, @@ -23,8 +23,13 @@ from vllm.entrypoints.openai.engine.protocol import ( ) from vllm.entrypoints.openai.models.serving import BaseModelPath, OpenAIServingModels from vllm.entrypoints.openai.parser.harmony_utils import get_encoding +from vllm.inputs import TokensPrompt from vllm.outputs import CompletionOutput, RequestOutput +from vllm.renderers.hf import HfRenderer +from vllm.renderers.mistral import MistralRenderer from vllm.tokenizers import get_tokenizer +from vllm.tokenizers.mistral import MistralTokenizer +from vllm.tokenizers.registry import tokenizer_args_from_config from vllm.tool_parsers import ToolParserManager from vllm.v1.engine.async_llm import AsyncLLM @@ -103,15 +108,16 @@ def gptoss_server(default_server_args: list[str]): @pytest.fixture(scope="class") def gptoss_speculative_server(default_server_args: list[str]): + attention_backend = ( + "TRITON_ATTN" + if not is_aiter_found_and_supported() + else "ROCM_AITER_UNIFIED_ATTN" + ) server_args = default_server_args + [ "--speculative-config", f'{{"model": "{GPT_OSS_SPECULATOR_NAME}", ' f'"method": "eagle3", "num_speculative_tokens": 3}}', - f"--attention-backend={ - 'TRITON_ATTN' - if not is_aiter_found_and_supported() - else 'ROCM_AITER_UNIFIED_ATTN' - }", + f"--attention-backend={attention_backend}", ] # gpt-oss requires AITER unified attention on ROCm # TODO: Remove after fixing TRITON_ATTN issue on ROCm @@ -520,12 +526,21 @@ class MockModelConfig: encoder_config = None generation_config: str = "auto" media_io_kwargs: dict[str, dict[str, Any]] = field(default_factory=dict) - skip_tokenizer_init = False + skip_tokenizer_init: bool = False def get_diff_sampling_param(self): return self.diff_sampling_param or {} +def _build_renderer(model_config: MockModelConfig): + _, tokenizer_name, _, kwargs = tokenizer_args_from_config(model_config) + + return HfRenderer( + model_config, + tokenizer_kwargs={**kwargs, "tokenizer_name": tokenizer_name}, + ) + + def _build_serving_chat(engine: AsyncLLM) -> OpenAIServingChat: models = OpenAIServingModels( engine_client=engine, @@ -561,6 +576,7 @@ class MockEngine: model_config: MockModelConfig = field(default_factory=MockModelConfig) input_processor: MagicMock = field(default_factory=MagicMock) io_processor: MagicMock = field(default_factory=MagicMock) + renderer: MagicMock = field(default_factory=MagicMock) async def _async_serving_chat_init(): @@ -586,11 +602,11 @@ def test_async_serving_chat_init(): @pytest.mark.asyncio async def test_serving_chat_returns_correct_model_name(): mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = MockModelConfig() mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) serving_chat = _build_serving_chat(mock_engine) messages = [{"role": "user", "content": "what is 1+1?"}] @@ -616,11 +632,11 @@ async def test_serving_chat_returns_correct_model_name(): @pytest.mark.asyncio async def test_serving_chat_should_set_correct_max_tokens(): mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = MockModelConfig() mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) serving_chat = _build_serving_chat(mock_engine) @@ -649,11 +665,11 @@ async def test_serving_chat_should_set_correct_max_tokens(): # Reinitialize the engine with new settings mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = mock_model_config mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) # Initialize the serving chat serving_chat = _build_serving_chat(mock_engine) @@ -694,11 +710,11 @@ async def test_serving_chat_should_set_correct_max_tokens(): # Reinitialize the engine with new settings mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = mock_model_config mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) # Initialize the serving chat serving_chat = _build_serving_chat(mock_engine) @@ -732,42 +748,32 @@ async def test_serving_chat_should_set_correct_max_tokens(): @pytest.mark.asyncio -async def test_serving_chat_mistral_token_ids_prompt_is_validated(monkeypatch_module): +async def test_serving_chat_mistral_token_ids_prompt_is_validated(): """Regression test: when the Mistral tokenizer path returns token IDs directly, we must still apply input length + max_tokens validation. """ mock_engine = MagicMock(spec=AsyncLLM) mock_engine.errored = False - mock_engine.model_config = MockModelConfig() + mock_engine.model_config = MockModelConfig(skip_tokenizer_init=True) mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() - class DummyMistralTokenizer: - def decode(self, token_ids): - # Only used for logging/validation error messages. - return "dummy" - - dummy_tokenizer = DummyMistralTokenizer() - mock_engine.get_tokenizer.return_value = dummy_tokenizer - - # Patch the OpenAI engine serving module to treat our dummy tokenizer - # as a MistralTokenizer. This forces the code path where chat template - # rendering can return a list[int] (token IDs). - import vllm.entrypoints.openai.engine.serving as engine_serving - - monkeypatch_module.setattr( - engine_serving, "MistralTokenizer", DummyMistralTokenizer - ) - - serving_chat = _build_serving_chat(mock_engine) - + mock_tokenizer = MagicMock(spec=MistralTokenizer) + mock_renderer = MistralRenderer(mock_engine.model_config, tokenizer_kwargs={}) + mock_renderer._tokenizer = mock_tokenizer # Force the Mistral chat template renderer to return token IDs. # Choose a prompt length that is < max_model_len, but large enough that # adding max_tokens should exceed the model context window. - serving_chat._apply_mistral_chat_template_async = AsyncMock( - return_value=list(range(95)) + mock_renderer.render_messages_async = AsyncMock( + return_value=( + [], + TokensPrompt(prompt_token_ids=list(range(95))), + ) ) + mock_engine.renderer = mock_renderer + + serving_chat = _build_serving_chat(mock_engine) req = ChatCompletionRequest( model=MODEL_NAME, @@ -781,39 +787,33 @@ async def test_serving_chat_mistral_token_ids_prompt_is_validated(monkeypatch_mo @pytest.mark.asyncio -async def test_serving_chat_mistral_token_ids_prompt_too_long_is_rejected( - monkeypatch_module, -): +async def test_serving_chat_mistral_token_ids_prompt_too_long_is_rejected(): """Regression test: MistralTokenizer token-id prompts must still enforce the max context length for the input itself (token_num >= max_model_len). """ mock_engine = MagicMock(spec=AsyncLLM) mock_engine.errored = False - mock_engine.model_config = MockModelConfig() + mock_engine.model_config = MockModelConfig(skip_tokenizer_init=True) mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() - class DummyMistralTokenizer: - def decode(self, token_ids): - return "dummy" - - dummy_tokenizer = DummyMistralTokenizer() - mock_engine.get_tokenizer.return_value = dummy_tokenizer - - import vllm.entrypoints.openai.engine.serving as engine_serving - - monkeypatch_module.setattr( - engine_serving, "MistralTokenizer", DummyMistralTokenizer - ) - - serving_chat = _build_serving_chat(mock_engine) - + mock_tokenizer = MagicMock(spec=MistralTokenizer) + mock_renderer = MistralRenderer(mock_engine.model_config, tokenizer_kwargs={}) + mock_renderer._tokenizer = mock_tokenizer # prompt_token_ids length == max_model_len should be rejected for # completion-like requests (ChatCompletionRequest). - serving_chat._apply_mistral_chat_template_async = AsyncMock( - return_value=list(range(mock_engine.model_config.max_model_len)) + mock_renderer.render_messages_async = AsyncMock( + return_value=( + [], + TokensPrompt( + prompt_token_ids=list(range(mock_engine.model_config.max_model_len)) + ), + ) ) + mock_engine.renderer = mock_renderer + + serving_chat = _build_serving_chat(mock_engine) req = ChatCompletionRequest( model=MODEL_NAME, @@ -835,11 +835,11 @@ async def test_serving_chat_could_load_correct_generation_config(): } mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = mock_model_config mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) # Initialize the serving chat serving_chat = _build_serving_chat(mock_engine) @@ -881,11 +881,11 @@ async def test_serving_chat_did_set_correct_cache_salt(model_type): mock_model_config.hf_config.model_type = model_type mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = mock_model_config mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) serving_chat = _build_serving_chat(mock_engine) @@ -914,11 +914,11 @@ async def test_serving_chat_data_parallel_rank_extraction(): """Test that data_parallel_rank is properly extracted from header and passed to engine.""" mock_engine = MagicMock(spec=AsyncLLM) - mock_engine.get_tokenizer.return_value = get_tokenizer(MODEL_NAME) mock_engine.errored = False mock_engine.model_config = MockModelConfig() mock_engine.input_processor = MagicMock() mock_engine.io_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) # Mock the generate method to return an async generator async def mock_generate(*args, **kwargs): diff --git a/tests/entrypoints/openai/test_serving_engine.py b/tests/entrypoints/openai/test_serving_engine.py deleted file mode 100644 index 654d42276a8..00000000000 --- a/tests/entrypoints/openai/test_serving_engine.py +++ /dev/null @@ -1,71 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -import asyncio -import time -from unittest.mock import Mock - -import pytest - -from vllm.config import ModelConfig -from vllm.entrypoints.openai.engine.serving import OpenAIServing -from vllm.entrypoints.openai.models.serving import OpenAIServingModels -from vllm.tokenizers.mistral import MistralTokenizer - - -@pytest.fixture() -def serving() -> OpenAIServing: - """Create a minimal OpenAIServing instance for testing.""" - - # Create minimal mocks - engine_client = Mock() - model_config = Mock(spec=ModelConfig) - model_config.max_model_len = 32768 - models = Mock(spec=OpenAIServingModels) - models.model_config = model_config - models.input_processor = Mock() - models.io_processor = Mock() - - serving = OpenAIServing( - engine_client=engine_client, - models=models, - request_logger=None, - ) - return serving - - -@pytest.mark.asyncio -async def test_async_mistral_tokenizer_does_not_block_event_loop( - serving: OpenAIServing, -): - expected_tokens = [1, 2, 3] - - # Mock the blocking version to sleep - def mocked_apply_chat_template(*_args, **_kwargs): - time.sleep(2) - return expected_tokens - - mock_tokenizer = Mock(spec=MistralTokenizer) - mock_tokenizer.apply_chat_template.side_effect = mocked_apply_chat_template - - task = serving._apply_mistral_chat_template_async( - tokenizer=mock_tokenizer, messages=[], chat_template=None, tools=[] - ) - - # Ensure the event loop is not blocked - blocked_count = 0 - for _i in range(20): # Check over ~2 seconds - start = time.perf_counter() - await asyncio.sleep(0) - elapsed = time.perf_counter() - start - - # an overly generous elapsed time for slow machines - if elapsed >= 0.5: - blocked_count += 1 - - await asyncio.sleep(0.1) - - # Ensure task completes - tokens = await task - assert tokens == expected_tokens, "Mocked blocking tokenizer was not called" - assert blocked_count == 0, "Event loop blocked during tokenization" diff --git a/tests/entrypoints/openai/test_serving_models.py b/tests/entrypoints/openai/test_serving_models.py index 88b168c7d1b..f6755f48934 100644 --- a/tests/entrypoints/openai/test_serving_models.py +++ b/tests/entrypoints/openai/test_serving_models.py @@ -35,6 +35,7 @@ async def _async_serving_models_init() -> OpenAIServingModels: mock_engine_client.model_config = mock_model_config mock_engine_client.input_processor = MagicMock() mock_engine_client.io_processor = MagicMock() + mock_engine_client.renderer = MagicMock() serving_models = OpenAIServingModels( engine_client=mock_engine_client, diff --git a/tests/entrypoints/openai/test_serving_responses.py b/tests/entrypoints/openai/test_serving_responses.py index e2c19c24c10..ba0c2c876e0 100644 --- a/tests/entrypoints/openai/test_serving_responses.py +++ b/tests/entrypoints/openai/test_serving_responses.py @@ -131,6 +131,7 @@ class TestInitializeToolSessions: engine_client.input_processor = MagicMock() engine_client.io_processor = MagicMock() + engine_client.renderer = MagicMock() models = MagicMock() @@ -217,6 +218,7 @@ class TestValidateGeneratorInput: engine_client.input_processor = MagicMock() engine_client.io_processor = MagicMock() + engine_client.renderer = MagicMock() models = MagicMock() diff --git a/tests/entrypoints/pooling/score/test_utils.py b/tests/entrypoints/pooling/score/test_utils.py index 0c8e567d085..d69da822dd0 100644 --- a/tests/entrypoints/pooling/score/test_utils.py +++ b/tests/entrypoints/pooling/score/test_utils.py @@ -212,7 +212,7 @@ class TestGetScorePrompt: return_value=mock_model_no_score_template, ), patch( - "vllm.entrypoints.pooling.score.utils.apply_hf_chat_template", + "vllm.entrypoints.pooling.score.utils.safe_apply_chat_template", return_value="test querytest doc", ), ): @@ -245,7 +245,7 @@ class TestGetScorePrompt: return_value=mock_model_no_score_template, ), patch( - "vllm.entrypoints.pooling.score.utils.apply_hf_chat_template", + "vllm.entrypoints.pooling.score.utils.safe_apply_chat_template", side_effect=ChatTemplateResolutionError("No template"), ), ): @@ -296,7 +296,7 @@ class TestGetScorePrompt: return_value=mock_model_no_score_template, ), patch( - "vllm.entrypoints.pooling.score.utils.apply_hf_chat_template", + "vllm.entrypoints.pooling.score.utils.safe_apply_chat_template", side_effect=ChatTemplateResolutionError("No template"), ), ): @@ -331,7 +331,7 @@ class TestGetScorePrompt: return_value=mock_model_with_score_template, ), patch( - "vllm.entrypoints.pooling.score.utils.apply_hf_chat_template", + "vllm.entrypoints.pooling.score.utils.safe_apply_chat_template", side_effect=ChatTemplateResolutionError("No template"), ), ): diff --git a/tests/entrypoints/test_chat_utils.py b/tests/entrypoints/test_chat_utils.py index 6df2d26f2f0..ba43f37bd27 100644 --- a/tests/entrypoints/test_chat_utils.py +++ b/tests/entrypoints/test_chat_utils.py @@ -7,21 +7,14 @@ from typing import Literal import pytest import torch -from mistral_common.tokens.tokenizers.base import SpecialTokenPolicy from vllm.assets.audio import AudioAsset from vllm.assets.image import ImageAsset from vllm.assets.video import VideoAsset from vllm.config import ModelConfig from vllm.entrypoints.chat_utils import ( - _try_extract_ast, - apply_mistral_chat_template, - load_chat_template, parse_chat_messages, - parse_chat_messages_futures, - resolve_chat_template_content_format, - resolve_chat_template_kwargs, - resolve_hf_chat_template, + parse_chat_messages_async, ) from vllm.multimodal import MultiModalDataDict, MultiModalUUIDDict from vllm.multimodal.utils import ( @@ -29,24 +22,11 @@ from vllm.multimodal.utils import ( encode_image_url, encode_video_url, ) -from vllm.tokenizers import get_tokenizer -from vllm.tokenizers.mistral import MistralTokenizer from vllm.utils.serial_utils import tensor2base64 -from ..models.registry import HF_EXAMPLE_MODELS -from ..utils import VLLM_PATH - -EXAMPLES_DIR = VLLM_PATH / "examples" - PHI3V_MODEL_ID = "microsoft/Phi-3.5-vision-instruct" -ULTRAVOX_MODEL_ID = "fixie-ai/ultravox-v0_5-llama-3_2-1b" QWEN2AUDIO_MODEL_ID = "Qwen/Qwen2-Audio-7B-Instruct" -QWEN2VL_MODEL_ID = "Qwen/Qwen2-VL-2B-Instruct" -QWEN25VL_MODEL_ID = "Qwen/Qwen2.5-VL-3B-Instruct" QWEN25OMNI_MODEL_ID = "Qwen/Qwen2.5-Omni-7B" -QWEN3_MODEL_ID = "Qwen/Qwen3-8B" -LLAMA_GUARD_MODEL_ID = "meta-llama/Llama-Guard-3-1B" -HERMES_MODEL_ID = "NousResearch/Hermes-3-Llama-3.1-8B" MISTRAL_MODEL_ID = "mistralai/Mistral-Small-3.1-24B-Instruct-2503" @@ -469,7 +449,7 @@ async def test_parse_chat_messages_single_image_with_uuid_async( image_url, ): image_uuid = str(hash(image_url)) - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -490,7 +470,7 @@ async def test_parse_chat_messages_single_image_with_uuid_async( assert conversation == [ {"role": "user", "content": "<|image_1|>\nWhat's in the image?"} ] - _assert_mm_data_is_image_input(await mm_future, 1) + _assert_mm_data_is_image_input(mm_data, 1) _assert_mm_uuids(mm_uuids, 1, expected_uuids=[image_uuid]) @@ -500,7 +480,7 @@ async def test_parse_chat_messages_empty_image_with_uuid_async( image_url, ): image_uuid = str(hash(image_url)) - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -521,7 +501,7 @@ async def test_parse_chat_messages_empty_image_with_uuid_async( assert conversation == [ {"role": "user", "content": "<|image_1|>\nWhat's in the image?"} ] - _assert_mm_data_is_image_input(await mm_future, 1, skipped_image_indices=[0]) + _assert_mm_data_is_image_input(mm_data, 1, skipped_image_indices=[0]) _assert_mm_uuids(mm_uuids, 1, expected_uuids=[image_uuid]) @@ -533,7 +513,7 @@ async def test_parse_chat_messages_multiple_images_with_uuids_async( image_uuid1 = "my_uuid_1" image_uuid2 = "my_uuid_2" - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -562,7 +542,7 @@ async def test_parse_chat_messages_multiple_images_with_uuids_async( "content": "<|image_1|>\n<|image_2|>\nWhat's in these images?", } ] - _assert_mm_data_is_image_input(await mm_future, 2) + _assert_mm_data_is_image_input(mm_data, 2) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[image_uuid1, image_uuid2]) @@ -574,7 +554,7 @@ async def test_parse_chat_messages_multiple_empty_images_with_uuids_async( image_uuid1 = "my_uuid_1" image_uuid2 = "my_uuid_2" - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -603,7 +583,7 @@ async def test_parse_chat_messages_multiple_empty_images_with_uuids_async( "content": "<|image_1|>\n<|image_2|>\nWhat's in these images?", } ] - _assert_mm_data_is_image_input(await mm_future, 2, skipped_image_indices=[0, 1]) + _assert_mm_data_is_image_input(mm_data, 2, skipped_image_indices=[0, 1]) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[image_uuid1, image_uuid2]) @@ -614,7 +594,7 @@ async def test_parse_chat_messages_multiple_images_with_partial_uuids_async( ): image_uuid2 = "my_uuid_2" - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -642,7 +622,7 @@ async def test_parse_chat_messages_multiple_images_with_partial_uuids_async( "content": "<|image_1|>\n<|image_2|>\nWhat's in these images?", } ] - _assert_mm_data_is_image_input(await mm_future, 2) + _assert_mm_data_is_image_input(mm_data, 2) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[None, image_uuid2]) @@ -689,7 +669,7 @@ async def test_parse_chat_messages_single_image_async( phi3v_model_config, image_url, ): - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -706,7 +686,7 @@ async def test_parse_chat_messages_single_image_async( assert conversation == [ {"role": "user", "content": "<|image_1|>\nWhat's in the image?"} ] - _assert_mm_data_is_image_input(await mm_future, 1) + _assert_mm_data_is_image_input(mm_data, 1) _assert_mm_uuids(mm_uuids, 1, expected_uuids=[None]) @@ -890,7 +870,7 @@ async def test_parse_chat_messages_audio_embeds_async( # Encode it as base64 base64_audio_embedding = tensor2base64(audio_embedding) - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -908,7 +888,6 @@ async def test_parse_chat_messages_audio_embeds_async( ) # Should have audio embedding in mm_data (single tensor, not a list) - mm_data = await mm_future assert mm_data is not None assert "audio" in mm_data assert isinstance(mm_data["audio"], torch.Tensor) @@ -1050,7 +1029,7 @@ async def test_parse_chat_messages_multiple_image_embeds_async( base64_image_embedding_1 = tensor2base64(image_embedding_1) base64_image_embedding_2 = tensor2base64(image_embedding_2) - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -1080,7 +1059,6 @@ async def test_parse_chat_messages_multiple_image_embeds_async( ] # Await the future and verify mm_data - mm_data = await mm_future assert mm_data is not None assert "image" in mm_data assert isinstance(mm_data["image"], list) @@ -1101,7 +1079,7 @@ async def test_parse_chat_messages_empty_image_embeds_with_uuid_async( phi3v_model_config_image_embeds, ): uuid = "abcd" - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -1121,7 +1099,6 @@ async def test_parse_chat_messages_empty_image_embeds_with_uuid_async( "content": "<|image_1|>\nWhat's in this image?", } ] - mm_data = await mm_future assert mm_data is not None assert "image" in mm_data assert isinstance(mm_data["image"], list) @@ -1228,7 +1205,7 @@ async def test_parse_chat_messages_multiple_images_async( phi3v_model_config, image_url, ): - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -1252,7 +1229,7 @@ async def test_parse_chat_messages_multiple_images_async( "content": "<|image_1|>\n<|image_2|>\nWhat's in these images?", } ] - _assert_mm_data_is_image_input(await mm_future, 2) + _assert_mm_data_is_image_input(mm_data, 2) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[None, None]) @@ -1582,7 +1559,7 @@ async def test_parse_chat_messages_multiple_images_interleave_async( phi3v_model_config_mm_interleaved, image_url, ): - conversation, mm_data, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -1609,7 +1586,7 @@ async def test_parse_chat_messages_multiple_images_interleave_async( "Do they have differences?", } ] - _assert_mm_data_is_image_input(await mm_data, 2) + _assert_mm_data_is_image_input(mm_data, 2) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[None, None]) @@ -1619,7 +1596,7 @@ async def test_parse_chat_messages_multiple_images_with_uuids_interleave_async( image_url, ): image_uuid = str(hash(image_url)) - conversation, mm_data, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -1654,7 +1631,7 @@ async def test_parse_chat_messages_multiple_images_with_uuids_interleave_async( "Do they have differences?", } ] - _assert_mm_data_is_image_input(await mm_data, 2) + _assert_mm_data_is_image_input(mm_data, 2) _assert_mm_uuids(mm_uuids, 2, expected_uuids=[image_uuid, image_uuid]) @@ -2030,377 +2007,6 @@ def test_parse_chat_messages_multiple_images_interleave_with_placeholders( ) -@pytest.mark.parametrize( - "model", - [ - QWEN2VL_MODEL_ID, # tokenizer.chat_template is of type str - HERMES_MODEL_ID, # tokenizer.chat_template is of type dict - ], -) -@pytest.mark.parametrize("use_tools", [True, False]) -def test_resolve_hf_chat_template(sample_json_schema, model, use_tools): - """checks that chat_template is a dict type for HF models.""" - model_info = HF_EXAMPLE_MODELS.find_hf_info(model) - model_info.check_available_online(on_fail="skip") - - model_config = ModelConfig( - model, - tokenizer=model_info.tokenizer or model, - tokenizer_mode=model_info.tokenizer_mode, - revision=model_info.revision, - trust_remote_code=model_info.trust_remote_code, - hf_overrides=model_info.hf_overrides, - skip_tokenizer_init=model_info.require_embed_inputs, - enable_prompt_embeds=model_info.require_embed_inputs, - enable_mm_embeds=model_info.require_embed_inputs, - enforce_eager=model_info.enforce_eager, - dtype=model_info.dtype, - ) - - # Build the tokenizer - tokenizer = get_tokenizer( - model, - trust_remote_code=model_config.trust_remote_code, - ) - - tools = ( - [ - { - "type": "function", - "function": { - "name": "dummy_function_name", - "description": "This is a dummy function", - "parameters": sample_json_schema, - }, - } - ] - if use_tools - else None - ) - - # Test detecting the tokenizer's chat_template - chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=None, - tools=tools, - model_config=model_config, - ) - assert isinstance(chat_template, str) - - -@pytest.mark.parametrize( - "model, expected_kwargs", - [ - ( - QWEN2VL_MODEL_ID, - { - "add_vision_id", - "add_generation_prompt", - "continue_final_message", - "tools", - }, - ), - ( - QWEN3_MODEL_ID, - { - "enable_thinking", - "add_generation_prompt", - "continue_final_message", - "tools", - }, - ), - ], -) -def test_resolve_hf_chat_template_kwargs(sample_json_schema, model, expected_kwargs): - """checks that chat_template is a dict type for HF models.""" - model_info = HF_EXAMPLE_MODELS.find_hf_info(model) - model_info.check_available_online(on_fail="skip") - - tools = [ - { - "type": "function", - "function": { - "name": "dummy_function_name", - "description": "This is a dummy function", - "parameters": sample_json_schema, - }, - } - ] - - chat_template_kwargs = { - # both unused - "unsed_kwargs_1": 123, - "unsed_kwargs_2": "abc", - # should not appear - "chat_template": "{% Hello world! %}", - "tokenize": True, - # used by tokenizer - "continue_final_message": True, - "tools": tools, - # both used by Qwen2-VL and Qwen3 - "add_generation_prompt": True, - # only used by Qwen2-VL - "add_vision_id": True, - # only used by Qwen3 - "enable_thinking": True, - } - - model_config = ModelConfig( - model, - tokenizer=model_info.tokenizer or model, - tokenizer_mode=model_info.tokenizer_mode, - revision=model_info.revision, - trust_remote_code=model_info.trust_remote_code, - hf_overrides=model_info.hf_overrides, - skip_tokenizer_init=model_info.require_embed_inputs, - enable_prompt_embeds=model_info.require_embed_inputs, - enable_mm_embeds=model_info.require_embed_inputs, - enforce_eager=model_info.enforce_eager, - dtype=model_info.dtype, - ) - - # Build the tokenizer - tokenizer = get_tokenizer( - model, - trust_remote_code=model_config.trust_remote_code, - ) - - # Test detecting the tokenizer's chat_template - chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=None, - tools=tools, - model_config=model_config, - ) - with pytest.raises( - ValueError, match="Found unexpected chat template kwargs from request" - ): - # should raise error if `chat_template_kwargs` contains - # `chat_template` or `tokenize` - resolve_chat_template_kwargs( - tokenizer, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - ) - resolved_chat_template_kwargs = resolve_chat_template_kwargs( - tokenizer, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - raise_on_unexpected=False, - ) - assert set(resolved_chat_template_kwargs.keys()) == expected_kwargs - - # Additional test: Verify HF base parameters work with **kwargs tokenizers - # This validates the fix for tokenizers like Kimi K2 that use **kwargs - # to receive standard HuggingFace parameters instead of declaring them explicitly - from vllm.entrypoints.chat_utils import _get_hf_base_chat_template_params - - hf_base_params = _get_hf_base_chat_template_params() - # Verify common HF parameters are in the base class - assert {"add_generation_prompt", "tools", "continue_final_message"}.issubset( - hf_base_params - ), f"Expected HF base params not found in {hf_base_params}" - - # Test with a mock tokenizer that uses **kwargs (like Kimi K2) - class MockTokenizerWithKwargs: - def apply_chat_template(self, conversation, **kwargs): - return "mocked_output" - - mock_tokenizer = MockTokenizerWithKwargs() - mock_kwargs = { - "add_generation_prompt": True, - "tools": tools, - "continue_final_message": False, - "unknown_param": "should_be_filtered", - } - resolved_mock = resolve_chat_template_kwargs( - mock_tokenizer, chat_template, mock_kwargs, raise_on_unexpected=False - ) - # HF base params should pass through even with **kwargs tokenizer - assert "add_generation_prompt" in resolved_mock - assert "tools" in resolved_mock - assert "continue_final_message" in resolved_mock - # Unknown params should be filtered out - assert "unknown_param" not in resolved_mock - - -# NOTE: Qwen2-Audio default chat template is specially defined inside -# processor class instead of using `tokenizer_config.json` -@pytest.mark.parametrize( - ("model", "expected_format"), - [ - (PHI3V_MODEL_ID, "string"), - (QWEN2VL_MODEL_ID, "openai"), - (QWEN25VL_MODEL_ID, "openai"), - (ULTRAVOX_MODEL_ID, "string"), - (QWEN2AUDIO_MODEL_ID, "openai"), - (LLAMA_GUARD_MODEL_ID, "openai"), - ], -) -def test_resolve_content_format_hf_defined(model, expected_format): - model_info = HF_EXAMPLE_MODELS.find_hf_info(model) - model_info.check_available_online(on_fail="skip") - - model_config = ModelConfig( - model, - tokenizer=model_info.tokenizer or model, - tokenizer_mode=model_info.tokenizer_mode, - revision=model_info.revision, - trust_remote_code=model_info.trust_remote_code, - hf_overrides=model_info.hf_overrides, - skip_tokenizer_init=model_info.require_embed_inputs, - enable_prompt_embeds=model_info.require_embed_inputs, - enable_mm_embeds=model_info.require_embed_inputs, - enforce_eager=model_info.enforce_eager, - dtype=model_info.dtype, - ) - - tokenizer = get_tokenizer( - model, - trust_remote_code=model_config.trust_remote_code, - ) - - # Test detecting the tokenizer's chat_template - chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=None, - tools=None, - model_config=model_config, - ) - assert isinstance(chat_template, str) - - print("[TEXT]") - print(chat_template) - print("[AST]") - print(_try_extract_ast(chat_template)) - - resolved_format = resolve_chat_template_content_format( - None, # Test detecting the tokenizer's chat_template - None, - "auto", - tokenizer, - model_config=model_config, - ) - - assert resolved_format == expected_format - - -@pytest.mark.parametrize( - ("model", "expected_format"), - [ - ("Salesforce/blip2-opt-2.7b", "string"), - ("facebook/chameleon-7b", "string"), - ("deepseek-ai/deepseek-vl2-tiny", "string"), - ("adept/fuyu-8b", "string"), - ("google/paligemma-3b-mix-224", "string"), - ("Qwen/Qwen-VL", "string"), - ("Qwen/Qwen-VL-Chat", "string"), - ], -) -def test_resolve_content_format_fallbacks(model, expected_format): - model_info = HF_EXAMPLE_MODELS.find_hf_info(model) - model_info.check_available_online(on_fail="skip") - - model_config = ModelConfig( - model, - tokenizer=model_info.tokenizer or model, - tokenizer_mode=model_info.tokenizer_mode, - revision=model_info.revision, - trust_remote_code=model_info.trust_remote_code, - hf_overrides=model_info.hf_overrides, - skip_tokenizer_init=model_info.require_embed_inputs, - enable_prompt_embeds=model_info.require_embed_inputs, - enable_mm_embeds=model_info.require_embed_inputs, - enforce_eager=model_info.enforce_eager, - dtype=model_info.dtype, - ) - - tokenizer = get_tokenizer( - model_config.tokenizer, - trust_remote_code=model_config.trust_remote_code, - ) - - # Test detecting the tokenizer's chat_template - chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=None, - tools=None, - model_config=model_config, - ) - assert isinstance(chat_template, str) - - print("[TEXT]") - print(chat_template) - print("[AST]") - print(_try_extract_ast(chat_template)) - - resolved_format = resolve_chat_template_content_format( - None, # Test detecting the tokenizer's chat_template - None, - "auto", - tokenizer, - model_config=model_config, - ) - - assert resolved_format == expected_format - - -@pytest.mark.parametrize( - ("template_path", "expected_format"), - [ - ("template_alpaca.jinja", "string"), - ("template_baichuan.jinja", "string"), - ("template_chatglm.jinja", "string"), - ("template_chatglm2.jinja", "string"), - ("template_chatml.jinja", "string"), - ("template_dse_qwen2_vl.jinja", "openai"), - ("template_falcon_180b.jinja", "string"), - ("template_falcon.jinja", "string"), - ("template_inkbot.jinja", "string"), - ("template_teleflm.jinja", "string"), - ("template_vlm2vec_phi3v.jinja", "openai"), - ("template_vlm2vec_qwen2vl.jinja", "openai"), - ("tool_chat_template_granite_20b_fc.jinja", "string"), - ("tool_chat_template_hermes.jinja", "string"), - ("tool_chat_template_internlm2_tool.jinja", "string"), - ("tool_chat_template_llama3.1_json.jinja", "openai"), - ("tool_chat_template_llama3.2_json.jinja", "openai"), - ("tool_chat_template_mistral_parallel.jinja", "string"), - ("tool_chat_template_mistral.jinja", "string"), - ], -) -def test_resolve_content_format_examples(template_path, expected_format): - model_config = ModelConfig( - PHI3V_MODEL_ID, # Dummy - tokenizer=PHI3V_MODEL_ID, # Dummy - trust_remote_code=True, - ) - - dummy_tokenizer = get_tokenizer( - PHI3V_MODEL_ID, # Dummy - trust_remote_code=model_config.trust_remote_code, - ) - dummy_tokenizer.chat_template = None - - chat_template = load_chat_template(EXAMPLES_DIR / template_path) - assert isinstance(chat_template, str) - - print("[TEXT]") - print(chat_template) - print("[AST]") - print(_try_extract_ast(chat_template)) - - resolved_format = resolve_chat_template_content_format( - chat_template, - None, - "auto", - dummy_tokenizer, - model_config=model_config, - ) - - assert resolved_format == expected_format - - def test_parse_chat_messages_include_thinking_chunk(mistral_model_config): messages = [ { @@ -2462,56 +2068,6 @@ def test_parse_chat_messages_include_thinking_chunk(mistral_model_config): assert conversation_with_thinking == expected_conversation -def test_apply_mistral_chat_template_thinking_chunk(): - messages = [ - { - "role": "system", - "content": [ - {"type": "text", "text": "You are a helpful assistant."}, - { - "type": "thinking", - "closed": True, - "thinking": "Only return the answer when you are confident.", - }, - ], - }, - {"role": "user", "content": "What is 2+2?"}, - { - "role": "assistant", - "content": [ - {"type": "text", "text": "Let me think about it."}, - {"type": "thinking", "closed": True, "thinking": "2+2 = 4"}, - { - "type": "text", - "text": "The answer is 4.", - }, - ], - }, - {"role": "user", "content": "Thanks, what is 3+3?"}, - ] - mistral_tokenizer = MistralTokenizer.from_pretrained( - "mistralai/Magistral-Small-2509" - ) - - tokens_ids = apply_mistral_chat_template( - mistral_tokenizer, messages, chat_template=None, tools=None - ) - - string_tokens = mistral_tokenizer.mistral.decode( - tokens_ids, special_token_policy=SpecialTokenPolicy.KEEP - ) - - expected_tokens = ( - r"[SYSTEM_PROMPT]You are a helpful assistant.[THINK]Only return the" - r" answer when you are confident.[/THINK][/SYSTEM_PROMPT]" - r"[INST]What is 2+2?[/INST]" - r"Let me think about it.[THINK]2+2 = 4[/THINK]The answer is 4." - r"[INST]Thanks, what is 3+3?[/INST]" - ) - - assert string_tokens == expected_tokens - - def test_parse_chat_messages_single_empty_audio_with_uuid( qwen2_audio_model_config, ): @@ -2550,7 +2106,7 @@ async def test_parse_chat_messages_single_empty_audio_with_uuid_async( qwen2_audio_model_config, ): audio_uuid = "abcd" - conversation, mm_future, mm_uuids = parse_chat_messages_futures( + conversation, mm_data, mm_uuids = await parse_chat_messages_async( [ { "role": "user", @@ -2575,5 +2131,5 @@ async def test_parse_chat_messages_single_empty_audio_with_uuid_async( "audio say?", } ] - _assert_mm_data_inputs(await mm_future, {"audio": 1}) + _assert_mm_data_inputs(mm_data, {"audio": 1}) _assert_mm_uuids(mm_uuids, 1, modality="audio", expected_uuids=[audio_uuid]) diff --git a/tests/renderers/__init__.py b/tests/renderers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/renderers/test_hf.py b/tests/renderers/test_hf.py new file mode 100644 index 00000000000..e262e1f555a --- /dev/null +++ b/tests/renderers/test_hf.py @@ -0,0 +1,537 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import load_chat_template +from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest +from vllm.renderers.hf import ( + _get_hf_base_chat_template_params, + _try_extract_ast, + resolve_chat_template, + resolve_chat_template_content_format, + resolve_chat_template_kwargs, + safe_apply_chat_template, +) +from vllm.tokenizers import get_tokenizer + +from ..models.registry import HF_EXAMPLE_MODELS +from ..utils import VLLM_PATH + +EXAMPLES_DIR = VLLM_PATH / "examples" + +chatml_jinja_path = VLLM_PATH / "examples/template_chatml.jinja" +assert chatml_jinja_path.exists() + +# Define models, templates, and their corresponding expected outputs +MODEL_TEMPLATE_GENERATION_OUTPUT = [ + ( + "facebook/opt-125m", + chatml_jinja_path, + True, + False, + """<|im_start|>user +Hello<|im_end|> +<|im_start|>assistant +Hi there!<|im_end|> +<|im_start|>user +What is the capital of<|im_end|> +<|im_start|>assistant +""", + ), + ( + "facebook/opt-125m", + chatml_jinja_path, + False, + False, + """<|im_start|>user +Hello<|im_end|> +<|im_start|>assistant +Hi there!<|im_end|> +<|im_start|>user +What is the capital of""", + ), + ( + "facebook/opt-125m", + chatml_jinja_path, + False, + True, + """<|im_start|>user +Hello<|im_end|> +<|im_start|>assistant +Hi there!<|im_end|> +<|im_start|>user +What is the capital of<|im_end|> +<|im_start|>assistant +The capital of""", + ), +] + +TEST_MESSAGES = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi there!"}, + {"role": "user", "content": "What is the capital of"}, +] +ASSISTANT_MESSAGE_TO_CONTINUE = {"role": "assistant", "content": "The capital of"} + + +def test_load_chat_template(): + # Testing chatml template + template_content = load_chat_template(chat_template=chatml_jinja_path) + + # Test assertions + assert template_content is not None + # Hard coded value for template_chatml.jinja + assert ( + template_content + == """{% for message in messages %}{{'<|im_start|>' + message['role'] + '\\n' + message['content']}}{% if (loop.last and add_generation_prompt) or not loop.last %}{{ '<|im_end|>' + '\\n'}}{% endif %}{% endfor %} +{% if add_generation_prompt and messages[-1]['role'] != 'assistant' %}{{ '<|im_start|>assistant\\n' }}{% endif %}""" # noqa: E501 + ) + + +def test_no_load_chat_template_filelike(): + # Testing chatml template + template = "../../examples/does_not_exist" + + with pytest.raises(ValueError, match="looks like a file path"): + load_chat_template(chat_template=template) + + +def test_no_load_chat_template_literallike(): + # Testing chatml template + template = "{{ messages }}" + + template_content = load_chat_template(chat_template=template) + + assert template_content == template + + +@pytest.mark.parametrize( + "model", + [ + "Qwen/Qwen2-VL-2B-Instruct", # chat_template is of type str + "NousResearch/Hermes-3-Llama-3.1-8B", # chat_template is of type dict + ], +) +@pytest.mark.parametrize("use_tools", [True, False]) +def test_resolve_chat_template(sample_json_schema, model, use_tools): + """checks that chat_template is a dict type for HF models.""" + model_info = HF_EXAMPLE_MODELS.find_hf_info(model) + model_info.check_available_online(on_fail="skip") + + model_config = ModelConfig( + model, + tokenizer=model_info.tokenizer or model, + tokenizer_mode=model_info.tokenizer_mode, + revision=model_info.revision, + trust_remote_code=model_info.trust_remote_code, + hf_overrides=model_info.hf_overrides, + skip_tokenizer_init=model_info.require_embed_inputs, + enable_prompt_embeds=model_info.require_embed_inputs, + enable_mm_embeds=model_info.require_embed_inputs, + enforce_eager=model_info.enforce_eager, + dtype=model_info.dtype, + ) + + # Build the tokenizer + tokenizer = get_tokenizer( + model, + trust_remote_code=model_config.trust_remote_code, + ) + + tools = ( + [ + { + "type": "function", + "function": { + "name": "dummy_function_name", + "description": "This is a dummy function", + "parameters": sample_json_schema, + }, + } + ] + if use_tools + else None + ) + + # Test detecting the tokenizer's chat_template + chat_template = resolve_chat_template( + tokenizer, + chat_template=None, + tools=tools, + model_config=model_config, + ) + assert isinstance(chat_template, str) + + +@pytest.mark.parametrize( + "model, expected_kwargs", + [ + ( + "Qwen/Qwen2-VL-2B-Instruct", + { + "add_vision_id", + "add_generation_prompt", + "continue_final_message", + "tools", + }, + ), + ( + "Qwen/Qwen3-8B", + { + "enable_thinking", + "add_generation_prompt", + "continue_final_message", + "tools", + }, + ), + ], +) +def test_resolve_chat_template_kwargs(sample_json_schema, model, expected_kwargs): + """checks that chat_template is a dict type for HF models.""" + model_info = HF_EXAMPLE_MODELS.find_hf_info(model) + model_info.check_available_online(on_fail="skip") + + tools = [ + { + "type": "function", + "function": { + "name": "dummy_function_name", + "description": "This is a dummy function", + "parameters": sample_json_schema, + }, + } + ] + + chat_template_kwargs = { + # both unused + "unsed_kwargs_1": 123, + "unsed_kwargs_2": "abc", + # should not appear + "chat_template": "{% Hello world! %}", + "tokenize": True, + # used by tokenizer + "continue_final_message": True, + "tools": tools, + # both used by Qwen2-VL and Qwen3 + "add_generation_prompt": True, + # only used by Qwen2-VL + "add_vision_id": True, + # only used by Qwen3 + "enable_thinking": True, + } + + model_config = ModelConfig( + model, + tokenizer=model_info.tokenizer or model, + tokenizer_mode=model_info.tokenizer_mode, + revision=model_info.revision, + trust_remote_code=model_info.trust_remote_code, + hf_overrides=model_info.hf_overrides, + skip_tokenizer_init=model_info.require_embed_inputs, + enable_prompt_embeds=model_info.require_embed_inputs, + enable_mm_embeds=model_info.require_embed_inputs, + enforce_eager=model_info.enforce_eager, + dtype=model_info.dtype, + ) + + # Build the tokenizer + tokenizer = get_tokenizer( + model, + trust_remote_code=model_config.trust_remote_code, + ) + + # Test detecting the tokenizer's chat_template + chat_template = resolve_chat_template( + tokenizer, + chat_template=None, + tools=tools, + model_config=model_config, + ) + with pytest.raises( + ValueError, match="Found unexpected chat template kwargs from request" + ): + # should raise error if `chat_template_kwargs` contains + # `chat_template` or `tokenize` + resolve_chat_template_kwargs( + tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + resolved_chat_template_kwargs = resolve_chat_template_kwargs( + tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + raise_on_unexpected=False, + ) + assert set(resolved_chat_template_kwargs.keys()) == expected_kwargs + + # Additional test: Verify HF base parameters work with **kwargs tokenizers + # This validates the fix for tokenizers like Kimi K2 that use **kwargs + # to receive standard HuggingFace parameters instead of declaring them explicitly + hf_base_params = _get_hf_base_chat_template_params() + # Verify common HF parameters are in the base class + assert {"add_generation_prompt", "tools", "continue_final_message"}.issubset( + hf_base_params + ), f"Expected HF base params not found in {hf_base_params}" + + # Test with a mock tokenizer that uses **kwargs (like Kimi K2) + class MockTokenizerWithKwargs: + def apply_chat_template(self, conversation, **kwargs): + return "mocked_output" + + mock_tokenizer = MockTokenizerWithKwargs() + mock_kwargs = { + "add_generation_prompt": True, + "tools": tools, + "continue_final_message": False, + "unknown_param": "should_be_filtered", + } + resolved_mock = resolve_chat_template_kwargs( + mock_tokenizer, chat_template, mock_kwargs, raise_on_unexpected=False + ) + # HF base params should pass through even with **kwargs tokenizer + assert "add_generation_prompt" in resolved_mock + assert "tools" in resolved_mock + assert "continue_final_message" in resolved_mock + # Unknown params should be filtered out + assert "unknown_param" not in resolved_mock + + +# NOTE: Qwen2-Audio default chat template is specially defined inside +# processor class instead of using `tokenizer_config.json` +@pytest.mark.parametrize( + ("model", "expected_format"), + [ + ("microsoft/Phi-3.5-vision-instruct", "string"), + ("Qwen/Qwen2-VL-2B-Instruct", "openai"), + ("Qwen/Qwen2.5-VL-3B-Instruct", "openai"), + ("fixie-ai/ultravox-v0_5-llama-3_2-1b", "string"), + ("Qwen/Qwen2-Audio-7B-Instruct", "openai"), + ("meta-llama/Llama-Guard-3-1B", "openai"), + ], +) +def test_resolve_content_format_hf_defined(model, expected_format): + model_info = HF_EXAMPLE_MODELS.find_hf_info(model) + model_info.check_available_online(on_fail="skip") + + model_config = ModelConfig( + model, + tokenizer=model_info.tokenizer or model, + tokenizer_mode=model_info.tokenizer_mode, + revision=model_info.revision, + trust_remote_code=model_info.trust_remote_code, + hf_overrides=model_info.hf_overrides, + skip_tokenizer_init=model_info.require_embed_inputs, + enable_prompt_embeds=model_info.require_embed_inputs, + enable_mm_embeds=model_info.require_embed_inputs, + enforce_eager=model_info.enforce_eager, + dtype=model_info.dtype, + ) + + tokenizer = get_tokenizer( + model, + trust_remote_code=model_config.trust_remote_code, + ) + + # Test detecting the tokenizer's chat_template + chat_template = resolve_chat_template( + tokenizer, + chat_template=None, + tools=None, + model_config=model_config, + ) + assert isinstance(chat_template, str) + + print("[TEXT]") + print(chat_template) + print("[AST]") + print(_try_extract_ast(chat_template)) + + resolved_format = resolve_chat_template_content_format( + None, # Test detecting the tokenizer's chat_template + None, + "auto", + tokenizer, + model_config=model_config, + ) + + assert resolved_format == expected_format + + +@pytest.mark.parametrize( + ("model", "expected_format"), + [ + ("Salesforce/blip2-opt-2.7b", "string"), + ("facebook/chameleon-7b", "string"), + ("deepseek-ai/deepseek-vl2-tiny", "string"), + ("adept/fuyu-8b", "string"), + ("google/paligemma-3b-mix-224", "string"), + ("Qwen/Qwen-VL", "string"), + ("Qwen/Qwen-VL-Chat", "string"), + ], +) +def test_resolve_content_format_fallbacks(model, expected_format): + model_info = HF_EXAMPLE_MODELS.find_hf_info(model) + model_info.check_available_online(on_fail="skip") + + model_config = ModelConfig( + model, + tokenizer=model_info.tokenizer or model, + tokenizer_mode=model_info.tokenizer_mode, + revision=model_info.revision, + trust_remote_code=model_info.trust_remote_code, + hf_overrides=model_info.hf_overrides, + skip_tokenizer_init=model_info.require_embed_inputs, + enable_prompt_embeds=model_info.require_embed_inputs, + enable_mm_embeds=model_info.require_embed_inputs, + enforce_eager=model_info.enforce_eager, + dtype=model_info.dtype, + ) + + tokenizer = get_tokenizer( + model_config.tokenizer, + trust_remote_code=model_config.trust_remote_code, + ) + + # Test detecting the tokenizer's chat_template + chat_template = resolve_chat_template( + tokenizer, + chat_template=None, + tools=None, + model_config=model_config, + ) + assert isinstance(chat_template, str) + + print("[TEXT]") + print(chat_template) + print("[AST]") + print(_try_extract_ast(chat_template)) + + resolved_format = resolve_chat_template_content_format( + None, # Test detecting the tokenizer's chat_template + None, + "auto", + tokenizer, + model_config=model_config, + ) + + assert resolved_format == expected_format + + +@pytest.mark.parametrize( + ("template_path", "expected_format"), + [ + ("template_alpaca.jinja", "string"), + ("template_baichuan.jinja", "string"), + ("template_chatglm.jinja", "string"), + ("template_chatglm2.jinja", "string"), + ("template_chatml.jinja", "string"), + ("template_dse_qwen2_vl.jinja", "openai"), + ("template_falcon_180b.jinja", "string"), + ("template_falcon.jinja", "string"), + ("template_inkbot.jinja", "string"), + ("template_teleflm.jinja", "string"), + ("template_vlm2vec_phi3v.jinja", "openai"), + ("template_vlm2vec_qwen2vl.jinja", "openai"), + ("tool_chat_template_granite_20b_fc.jinja", "string"), + ("tool_chat_template_hermes.jinja", "string"), + ("tool_chat_template_internlm2_tool.jinja", "string"), + ("tool_chat_template_llama3.1_json.jinja", "openai"), + ("tool_chat_template_llama3.2_json.jinja", "openai"), + ("tool_chat_template_mistral_parallel.jinja", "string"), + ("tool_chat_template_mistral.jinja", "string"), + ], +) +def test_resolve_content_format_examples(template_path, expected_format): + model = "Qwen/Qwen2-VL-2B-Instruct" # Dummy + model_config = ModelConfig( + model, + tokenizer=model, + trust_remote_code=True, + ) + + dummy_tokenizer = get_tokenizer( + model, + trust_remote_code=model_config.trust_remote_code, + ) + dummy_tokenizer.chat_template = None + + chat_template = load_chat_template(EXAMPLES_DIR / template_path) + assert isinstance(chat_template, str) + + print("[TEXT]") + print(chat_template) + print("[AST]") + print(_try_extract_ast(chat_template)) + + resolved_format = resolve_chat_template_content_format( + chat_template, + None, + "auto", + dummy_tokenizer, + model_config=model_config, + ) + + assert resolved_format == expected_format + + +@pytest.mark.parametrize( + "model,template,add_generation_prompt,continue_final_message,expected_output", + MODEL_TEMPLATE_GENERATION_OUTPUT, +) +def test_get_gen_prompt( + model, template, add_generation_prompt, continue_final_message, expected_output +): + model_info = HF_EXAMPLE_MODELS.find_hf_info(model) + model_info.check_available_online(on_fail="skip") + + model_config = ModelConfig( + model, + tokenizer=model_info.tokenizer or model, + tokenizer_mode=model_info.tokenizer_mode, + trust_remote_code=model_info.trust_remote_code, + revision=model_info.revision, + hf_overrides=model_info.hf_overrides, + skip_tokenizer_init=model_info.require_embed_inputs, + enable_prompt_embeds=model_info.require_embed_inputs, + enable_mm_embeds=model_info.require_embed_inputs, + enforce_eager=model_info.enforce_eager, + dtype=model_info.dtype, + ) + + # Initialize the tokenizer + tokenizer = get_tokenizer( + tokenizer_name=model_config.tokenizer, + trust_remote_code=model_config.trust_remote_code, + ) + template_content = load_chat_template(chat_template=template) + + # Create a mock request object using keyword arguments + mock_request = ChatCompletionRequest( + model=model, + messages=TEST_MESSAGES + [ASSISTANT_MESSAGE_TO_CONTINUE] + if continue_final_message + else TEST_MESSAGES, + add_generation_prompt=add_generation_prompt, + continue_final_message=continue_final_message, + ) + + # Call the function and get the result + result = safe_apply_chat_template( + model_config, + tokenizer, + mock_request.messages, + tools=None, + chat_template=mock_request.chat_template or template_content, + add_generation_prompt=mock_request.add_generation_prompt, + continue_final_message=mock_request.continue_final_message, + tokenize=False, + ) + + # Test assertion + assert result == expected_output, ( + f"The generated prompt does not match the expected output for " + f"model {model} and template {template}" + ) diff --git a/tests/renderers/test_mistral.py b/tests/renderers/test_mistral.py new file mode 100644 index 00000000000..0dc214ae939 --- /dev/null +++ b/tests/renderers/test_mistral.py @@ -0,0 +1,100 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import asyncio +import time +from unittest.mock import Mock + +import pytest +from mistral_common.tokens.tokenizers.base import SpecialTokenPolicy + +from vllm.config import ModelConfig +from vllm.renderers.mistral import MistralRenderer, safe_apply_chat_template +from vllm.tokenizers.mistral import MistralTokenizer + + +@pytest.mark.asyncio +async def test_async_mistral_tokenizer_does_not_block_event_loop(): + expected_tokens = [1, 2, 3] + + # Mock the blocking version to sleep + def mocked_apply_chat_template(*_args, **_kwargs): + time.sleep(2) + return expected_tokens + + mock_tokenizer = Mock(spec=MistralTokenizer) + mock_tokenizer.apply_chat_template = mocked_apply_chat_template + mock_renderer = MistralRenderer(Mock(spec=ModelConfig), tokenizer_kwargs={}) + mock_renderer._tokenizer = mock_tokenizer + + task = mock_renderer.render_messages_async([]) + + # Ensure the event loop is not blocked + blocked_count = 0 + for _i in range(20): # Check over ~2 seconds + start = time.perf_counter() + await asyncio.sleep(0) + elapsed = time.perf_counter() - start + + # an overly generous elapsed time for slow machines + if elapsed >= 0.5: + blocked_count += 1 + + await asyncio.sleep(0.1) + + # Ensure task completes + _, prompt = await task + assert prompt["prompt_token_ids"] == expected_tokens, ( + "Mocked blocking tokenizer was not called" + ) + assert blocked_count == 0, "Event loop blocked during tokenization" + + +def test_apply_mistral_chat_template_thinking_chunk(): + messages = [ + { + "role": "system", + "content": [ + {"type": "text", "text": "You are a helpful assistant."}, + { + "type": "thinking", + "closed": True, + "thinking": "Only return the answer when you are confident.", + }, + ], + }, + {"role": "user", "content": "What is 2+2?"}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me think about it."}, + {"type": "thinking", "closed": True, "thinking": "2+2 = 4"}, + { + "type": "text", + "text": "The answer is 4.", + }, + ], + }, + {"role": "user", "content": "Thanks, what is 3+3?"}, + ] + mistral_tokenizer = MistralTokenizer.from_pretrained( + "mistralai/Magistral-Small-2509" + ) + + tokens_ids = safe_apply_chat_template( + mistral_tokenizer, messages, chat_template=None, tools=None + ) + + string_tokens = mistral_tokenizer.mistral.decode( + tokens_ids, special_token_policy=SpecialTokenPolicy.KEEP + ) + + expected_tokens = ( + r"[SYSTEM_PROMPT]You are a helpful assistant.[THINK]Only return the" + r" answer when you are confident.[/THINK][/SYSTEM_PROMPT]" + r"[INST]What is 2+2?[/INST]" + r"Let me think about it.[THINK]2+2 = 4[/THINK]The answer is 4." + r"[INST]Thanks, what is 3+3?[/INST]" + ) + + assert string_tokens == expected_tokens diff --git a/tests/test_inputs.py b/tests/test_inputs.py index 073be24a4a0..6ea4f465cdf 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -7,7 +7,6 @@ from vllm.config import ModelConfig from vllm.inputs import zip_enc_dec_prompts from vllm.inputs.parse import parse_raw_prompts from vllm.inputs.preprocess import InputPreprocessor -from vllm.tokenizers import cached_tokenizer_from_config pytestmark = pytest.mark.cpu_test @@ -115,10 +114,10 @@ def test_zip_enc_dec_prompts(mm_processor_kwargs, expected_mm_kwargs): ) def test_preprocessor_always_mm_code_path(model_id, prompt): model_config = ModelConfig(model=model_id) - tokenizer = cached_tokenizer_from_config(model_config) - input_preprocessor = InputPreprocessor(model_config, tokenizer) + input_preprocessor = InputPreprocessor(model_config) # HF processor adds sep token + tokenizer = input_preprocessor.get_tokenizer() sep_token_id = tokenizer.vocab[tokenizer.sep_token] processed_inputs = input_preprocessor.preprocess(prompt) diff --git a/tests/v1/engine/test_llm_engine.py b/tests/v1/engine/test_llm_engine.py index c1d5f8af791..7e5196efc87 100644 --- a/tests/v1/engine/test_llm_engine.py +++ b/tests/v1/engine/test_llm_engine.py @@ -224,7 +224,7 @@ def test_skip_tokenizer_initialization(model: str): ) sampling_params = SamplingParams(prompt_logprobs=True, detokenize=True) - with pytest.raises(ValueError, match="cannot pass text prompts when"): + with pytest.raises(ValueError, match="`skip_tokenizer_init=True`"): llm.generate("abc", sampling_params) outputs = llm.generate( diff --git a/tests/v1/engine/test_process_multi_modal_uuids.py b/tests/v1/engine/test_process_multi_modal_uuids.py index dbf9ffd2942..998ae89f861 100644 --- a/tests/v1/engine/test_process_multi_modal_uuids.py +++ b/tests/v1/engine/test_process_multi_modal_uuids.py @@ -5,7 +5,13 @@ import pytest from vllm.assets.image import ImageAsset from vllm.assets.video import VideoAsset -from vllm.config import CacheConfig, DeviceConfig, ModelConfig, VllmConfig +from vllm.config import ( + CacheConfig, + DeviceConfig, + ModelConfig, + MultiModalConfig, + VllmConfig, +) from vllm.multimodal import MultiModalRegistry, MultiModalUUIDDict from vllm.sampling_params import SamplingParams from vllm.v1.engine.input_processor import InputProcessor @@ -44,27 +50,22 @@ def _mock_input_processor( monkeypatch.setattr(VllmConfig, "__post_init__", lambda self: None, raising=True) model_config = ModelConfig( + tokenizer="dummy", skip_tokenizer_init=True, max_model_len=128, mm_processor_cache_gb=mm_cache_gb, generation_config="vllm", - tokenizer="dummy", ) + model_config.runner_type = "generate" + model_config.multimodal_config = MultiModalConfig(mm_processor_cache_gb=mm_cache_gb) - # Minimal multimodal_config to satisfy references in - # Processor.process_inputs. - class _MockMMConfig: - def __init__(self, gb: float): - self.mm_processor_cache_gb = gb - - model_config.multimodal_config = _MockMMConfig(mm_cache_gb) # type: ignore[attr-defined] vllm_config = VllmConfig( model_config=model_config, cache_config=CacheConfig(enable_prefix_caching=enable_prefix_caching), device_config=DeviceConfig(device="cpu"), ) - return InputProcessor(vllm_config, tokenizer=None) + return InputProcessor(vllm_config) def test_multi_modal_uuids_length_mismatch_raises(monkeypatch): diff --git a/tools/pre_commit/mypy.py b/tools/pre_commit/mypy.py index 48803930d7b..4cda869cae3 100755 --- a/tools/pre_commit/mypy.py +++ b/tools/pre_commit/mypy.py @@ -35,6 +35,7 @@ FILES = [ "vllm/multimodal", "vllm/platforms", "vllm/plugins", + "vllm/renderers", "vllm/tokenizers", "vllm/transformers_utils", "vllm/triton_utils", diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index bf656cf23de..205efd1d582 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -11,9 +11,9 @@ from vllm.lora.request import LoRARequest from vllm.outputs import PoolingRequestOutput, RequestOutput from vllm.plugins.io_processors import IOProcessor from vllm.pooling_params import PoolingParams +from vllm.renderers import RendererLike from vllm.sampling_params import SamplingParams from vllm.tasks import SupportedTask -from vllm.tokenizers import TokenizerLike from vllm.v1.engine import EngineCoreRequest from vllm.v1.engine.input_processor import InputProcessor @@ -26,6 +26,10 @@ class EngineClient(ABC): input_processor: InputProcessor io_processor: IOProcessor | None + @property + @abstractmethod + def renderer(self) -> RendererLike: ... + @property @abstractmethod def is_running(self) -> bool: ... @@ -88,11 +92,6 @@ class EngineClient(ABC): """ ... - @abstractmethod - async def get_tokenizer(self) -> TokenizerLike: - """Get the tokenizer""" - ... - @abstractmethod async def is_tracing_enabled(self) -> bool: ... diff --git a/vllm/entrypoints/chat_utils.py b/vllm/entrypoints/chat_utils.py index 5e31f60ad0c..eb796c9661f 100644 --- a/vllm/entrypoints/chat_utils.py +++ b/vllm/entrypoints/chat_utils.py @@ -2,22 +2,15 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import asyncio -import inspect import json +import warnings from abc import ABC, abstractmethod -from collections import Counter, defaultdict, deque +from collections import Counter, defaultdict from collections.abc import Awaitable, Callable, Iterable from functools import cached_property, lru_cache, partial from pathlib import Path from typing import TYPE_CHECKING, Any, Generic, Literal, TypeAlias, TypeVar, cast -import jinja2 -import jinja2.ext -import jinja2.meta -import jinja2.nodes -import jinja2.parser -import jinja2.sandbox -import transformers.utils.chat_template_utils as hf_chat_utils from openai.types.chat import ( ChatCompletionAssistantMessageParam, ChatCompletionContentPartImageParam, @@ -39,7 +32,6 @@ from openai.types.responses import ResponseInputImageParam from openai_harmony import Message as OpenAIHarmonyMessage from PIL import Image from pydantic import BaseModel, ConfigDict, TypeAdapter -from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast, ProcessorMixin # pydantic needs the TypedDict from typing_extensions from typing_extensions import Required, TypedDict @@ -50,24 +42,35 @@ from vllm.logger import init_logger from vllm.model_executor.models import SupportsMultiModal from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalDataDict, MultiModalUUIDDict from vllm.multimodal.utils import MEDIA_CONNECTOR_REGISTRY, MediaConnector -from vllm.tokenizers import TokenizerLike -from vllm.transformers_utils.chat_templates import get_chat_template_fallback_path -from vllm.transformers_utils.processor import cached_get_processor from vllm.utils import random_uuid from vllm.utils.collection_utils import is_list_of -from vllm.utils.func_utils import supports_kw from vllm.utils.import_utils import LazyLoader if TYPE_CHECKING: import torch - - from vllm.tokenizers.mistral import MistralTokenizer else: torch = LazyLoader("torch", globals(), "torch") logger = init_logger(__name__) +def __getattr__(name: str): + if name == "resolve_hf_chat_template": + from vllm.renderers.hf import resolve_chat_template + + warnings.warn( + "`vllm.entrypoints.chat_utils.resolve_hf_chat_template` has been moved to " + "`vllm.renderers.hf.resolve_chat_template`. " + "The old name will be removed in v0.16.", + DeprecationWarning, + stacklevel=2, + ) + + return resolve_chat_template + + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + class ChatTemplateResolutionError(ValueError): """Raised when chat template resolution fails. @@ -320,325 +323,8 @@ class ConversationMessage(TypedDict, total=False): # Passed in by user ChatTemplateContentFormatOption = Literal["auto", "string", "openai"] -# Used internally -_ChatTemplateContentFormat = Literal["string", "openai"] - - -def _is_var_access(node: jinja2.nodes.Node, varname: str) -> bool: - if isinstance(node, jinja2.nodes.Name): - return node.ctx == "load" and node.name == varname - - return False - - -def _is_attr_access(node: jinja2.nodes.Node, varname: str, key: str) -> bool: - if isinstance(node, jinja2.nodes.Getitem): - return ( - _is_var_access(node.node, varname) - and isinstance(node.arg, jinja2.nodes.Const) - and node.arg.value == key - ) - - if isinstance(node, jinja2.nodes.Getattr): - return _is_var_access(node.node, varname) and node.attr == key - - return False - - -def _is_var_or_elems_access( - node: jinja2.nodes.Node, - varname: str, - key: str | None = None, -) -> bool: - if isinstance(node, jinja2.nodes.Filter): - return node.node is not None and _is_var_or_elems_access( - node.node, varname, key - ) - if isinstance(node, jinja2.nodes.Test): - return _is_var_or_elems_access(node.node, varname, key) - - if isinstance(node, jinja2.nodes.Getitem) and isinstance( - node.arg, jinja2.nodes.Slice - ): - return _is_var_or_elems_access(node.node, varname, key) - - return _is_attr_access(node, varname, key) if key else _is_var_access(node, varname) - - -def _iter_nodes_assign_var_or_elems(root: jinja2.nodes.Node, varname: str): - # Global variable that is implicitly defined at the root - yield root, varname - - # Iterative BFS - related_varnames = deque([varname]) - while related_varnames: - related_varname = related_varnames.popleft() - - for assign_ast in root.find_all(jinja2.nodes.Assign): - lhs = assign_ast.target - rhs = assign_ast.node - - if _is_var_or_elems_access(rhs, related_varname): - assert isinstance(lhs, jinja2.nodes.Name) - yield assign_ast, lhs.name - - # Avoid infinite looping for self-assignment - if lhs.name != related_varname: - related_varnames.append(lhs.name) - - -# NOTE: The proper way to handle this is to build a CFG so that we can handle -# the scope in which each variable is defined, but that is too complicated -def _iter_nodes_assign_messages_item(root: jinja2.nodes.Node): - messages_varnames = [ - varname for _, varname in _iter_nodes_assign_var_or_elems(root, "messages") - ] - - # Search for {%- for message in messages -%} loops - for loop_ast in root.find_all(jinja2.nodes.For): - loop_iter = loop_ast.iter - loop_target = loop_ast.target - - for varname in messages_varnames: - if _is_var_or_elems_access(loop_iter, varname): - assert isinstance(loop_target, jinja2.nodes.Name) - yield loop_ast, loop_target.name - break - - -def _iter_nodes_assign_content_item(root: jinja2.nodes.Node): - message_varnames = [ - varname for _, varname in _iter_nodes_assign_messages_item(root) - ] - - # Search for {%- for content in message['content'] -%} loops - for loop_ast in root.find_all(jinja2.nodes.For): - loop_iter = loop_ast.iter - loop_target = loop_ast.target - - for varname in message_varnames: - if _is_var_or_elems_access(loop_iter, varname, "content"): - assert isinstance(loop_target, jinja2.nodes.Name) - yield loop_ast, loop_target.name - break - - -def _try_extract_ast(chat_template: str) -> jinja2.nodes.Template | None: - try: - jinja_compiled = hf_chat_utils._compile_jinja_template(chat_template) - return jinja_compiled.environment.parse(chat_template) - except Exception: - logger.exception("Error when compiling Jinja template") - return None - - -@lru_cache(maxsize=32) -def _detect_content_format( - chat_template: str, - *, - default: _ChatTemplateContentFormat, -) -> _ChatTemplateContentFormat: - jinja_ast = _try_extract_ast(chat_template) - if jinja_ast is None: - return default - - try: - next(_iter_nodes_assign_content_item(jinja_ast)) - except StopIteration: - return "string" - except Exception: - logger.exception("Error when parsing AST of Jinja template") - return default - else: - return "openai" - - -def resolve_mistral_chat_template( - chat_template: str | None, - **kwargs: Any, -) -> str | None: - if chat_template is not None or kwargs.get("chat_template_kwargs") is not None: - raise ValueError( - "'chat_template' or 'chat_template_kwargs' cannot be overridden " - "for mistral tokenizer." - ) - - return None - - -_PROCESSOR_CHAT_TEMPLATES = dict[tuple[str, bool], str | None]() -""" -Used in `_try_get_processor_chat_template` to avoid calling -`cached_get_processor` again if the processor fails to be loaded. - -This is needed because `lru_cache` does not cache when an exception happens. -""" - - -def _try_get_processor_chat_template( - tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast, - model_config: ModelConfig, -) -> str | None: - cache_key = (tokenizer.name_or_path, model_config.trust_remote_code) - if cache_key in _PROCESSOR_CHAT_TEMPLATES: - return _PROCESSOR_CHAT_TEMPLATES[cache_key] - - try: - processor = cached_get_processor( - tokenizer.name_or_path, - processor_cls=( - PreTrainedTokenizer, - PreTrainedTokenizerFast, - ProcessorMixin, - ), - trust_remote_code=model_config.trust_remote_code, - ) - if ( - isinstance(processor, ProcessorMixin) - and hasattr(processor, "chat_template") - and (chat_template := processor.chat_template) is not None - ): - _PROCESSOR_CHAT_TEMPLATES[cache_key] = chat_template - return chat_template - except Exception: - logger.debug( - "Failed to load AutoProcessor chat template for %s", - tokenizer.name_or_path, - exc_info=True, - ) - - _PROCESSOR_CHAT_TEMPLATES[cache_key] = None - return None - - -def resolve_hf_chat_template( - tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast, - chat_template: str | None, - tools: list[dict[str, Any]] | None, - *, - model_config: ModelConfig, -) -> str | None: - # 1st priority: The given chat template - if chat_template is not None: - return chat_template - - # 2nd priority: AutoProcessor chat template, unless tool calling is enabled - if tools is None: - chat_template = _try_get_processor_chat_template(tokenizer, model_config) - if chat_template is not None: - return chat_template - - # 3rd priority: AutoTokenizer chat template - try: - return tokenizer.get_chat_template(chat_template, tools=tools) - except Exception: - logger.debug( - "Failed to load AutoTokenizer chat template for %s", - tokenizer.name_or_path, - exc_info=True, - ) - - # 4th priority: Predefined fallbacks - path = get_chat_template_fallback_path( - model_type=model_config.hf_config.model_type, - tokenizer_name_or_path=model_config.tokenizer, - ) - if path is not None: - logger.info_once( - "Loading chat template fallback for %s as there isn't one " - "defined on HF Hub.", - tokenizer.name_or_path, - ) - chat_template = load_chat_template(path) - else: - logger.debug_once( - "There is no chat template fallback for %s", tokenizer.name_or_path - ) - - return chat_template - - -def _resolve_chat_template_content_format( - chat_template: str | None, - tools: list[dict[str, Any]] | None, - tokenizer: TokenizerLike | None, - *, - model_config: ModelConfig, -) -> _ChatTemplateContentFormat: - if isinstance(tokenizer, (PreTrainedTokenizer, PreTrainedTokenizerFast)): - hf_chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=chat_template, - tools=tools, - model_config=model_config, - ) - else: - hf_chat_template = None - - jinja_text = ( - hf_chat_template - if isinstance(hf_chat_template, str) - else load_chat_template(chat_template, is_literal=True) - ) - - detected_format = ( - "string" - if jinja_text is None - else _detect_content_format(jinja_text, default="string") - ) - - return detected_format - - -@lru_cache -def _log_chat_template_content_format( - chat_template: str | None, - given_format: ChatTemplateContentFormatOption, - detected_format: ChatTemplateContentFormatOption, -): - logger.info( - "Detected the chat template content format to be '%s'. " - "You can set `--chat-template-content-format` to override this.", - detected_format, - ) - - if given_format != "auto" and given_format != detected_format: - logger.warning( - "You specified `--chat-template-content-format %s` " - "which is different from the detected format '%s'. " - "If our automatic detection is incorrect, please consider " - "opening a GitHub issue so that we can improve it: " - "https://github.com/vllm-project/vllm/issues/new/choose", - given_format, - detected_format, - ) - - -def resolve_chat_template_content_format( - chat_template: str | None, - tools: list[dict[str, Any]] | None, - given_format: ChatTemplateContentFormatOption, - tokenizer: TokenizerLike | None, - *, - model_config: ModelConfig, -) -> _ChatTemplateContentFormat: - if given_format != "auto": - return given_format - - detected_format = _resolve_chat_template_content_format( - chat_template, - tools, - tokenizer, - model_config=model_config, - ) - - _log_chat_template_content_format( - chat_template, - given_format=given_format, - detected_format=detected_format, - ) - - return detected_format +# After resolving "auto" +ChatTemplateContentFormat = Literal["string", "openai"] ModalityStr = Literal["image", "audio", "video", "image_embeds", "audio_embeds"] @@ -1593,7 +1279,7 @@ _ToolParser = partial(cast, ChatCompletionToolMessageParam) def _parse_chat_message_content( message: ChatCompletionMessageParam, mm_tracker: BaseMultiModalItemTracker, - content_format: _ChatTemplateContentFormat, + content_format: ChatTemplateContentFormat, interleave_strings: bool, ) -> list[ConversationMessage]: role = message["role"] @@ -1669,7 +1355,7 @@ def _postprocess_messages(messages: list[ConversationMessage]) -> None: def parse_chat_messages( messages: list[ChatCompletionMessageParam], model_config: ModelConfig, - content_format: _ChatTemplateContentFormat, + content_format: ChatTemplateContentFormat, ) -> tuple[ list[ConversationMessage], MultiModalDataDict | None, @@ -1697,13 +1383,13 @@ def parse_chat_messages( return conversation, mm_tracker.all_mm_data(), mm_tracker.all_mm_uuids() -def parse_chat_messages_futures( +async def parse_chat_messages_async( messages: list[ChatCompletionMessageParam], model_config: ModelConfig, - content_format: _ChatTemplateContentFormat, + content_format: ChatTemplateContentFormat, ) -> tuple[ list[ConversationMessage], - Awaitable[MultiModalDataDict | None], + MultiModalDataDict | None, MultiModalUUIDDict | None, ]: conversation: list[ConversationMessage] = [] @@ -1725,174 +1411,7 @@ def parse_chat_messages_futures( _postprocess_messages(conversation) - return conversation, mm_tracker.all_mm_data(), mm_tracker.all_mm_uuids() - - -# adapted from https://github.com/huggingface/transformers/blob/v4.56.2/src/transformers/utils/chat_template_utils.py#L398-L412 -# only preserve the parse function used to resolve chat template kwargs -class AssistantTracker(jinja2.ext.Extension): - tags = {"generation"} - - def parse(self, parser: jinja2.parser.Parser) -> jinja2.nodes.CallBlock: - lineno = next(parser.stream).lineno - body = parser.parse_statements(["name:endgeneration"], drop_needle=True) - call = self.call_method("_generation_support") - call_block = jinja2.nodes.CallBlock(call, [], [], body) - return call_block.set_lineno(lineno) - - -def _resolve_chat_template_kwargs( - chat_template: str, -): - env = jinja2.sandbox.ImmutableSandboxedEnvironment( - trim_blocks=True, - lstrip_blocks=True, - extensions=[AssistantTracker, jinja2.ext.loopcontrols], - ) - parsed_content = env.parse(chat_template) - template_vars = jinja2.meta.find_undeclared_variables(parsed_content) - return template_vars - - -_cached_resolve_chat_template_kwargs = lru_cache(_resolve_chat_template_kwargs) - - -@lru_cache -def _get_hf_base_chat_template_params() -> frozenset[str]: - # Get standard parameters from HuggingFace's base tokenizer class. - # This dynamically extracts parameters from PreTrainedTokenizer's - # apply_chat_template method, ensuring compatibility with tokenizers - # that use **kwargs to receive standard parameters. - - # Read signature from HF's base class - the single source of truth - base_sig = inspect.signature(PreTrainedTokenizer.apply_chat_template) - # Exclude VAR_KEYWORD (**kwargs) and VAR_POSITIONAL (*args) placeholders - return frozenset( - p.name - for p in base_sig.parameters.values() - if p.kind - not in (inspect.Parameter.VAR_KEYWORD, inspect.Parameter.VAR_POSITIONAL) - ) - - -def resolve_chat_template_kwargs( - tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast, - chat_template: str, - chat_template_kwargs: dict[str, Any], - raise_on_unexpected: bool = True, -) -> dict[str, Any]: - # We exclude chat_template from kwargs here, because - # chat template has been already resolved at this stage - unexpected_vars = {"chat_template", "tokenize"} - if raise_on_unexpected and ( - unexpected_in_kwargs := unexpected_vars & chat_template_kwargs.keys() - ): - raise ValueError( - "Found unexpected chat template kwargs from request: " - f"{unexpected_in_kwargs}" - ) - - fn_kw = { - k - for k in chat_template_kwargs - if supports_kw(tokenizer.apply_chat_template, k, allow_var_kwargs=False) - } - template_vars = _cached_resolve_chat_template_kwargs(chat_template) - - # Allow standard HF parameters even if tokenizer uses **kwargs to receive them - hf_base_params = _get_hf_base_chat_template_params() - - accept_vars = (fn_kw | template_vars | hf_base_params) - unexpected_vars - return {k: v for k, v in chat_template_kwargs.items() if k in accept_vars} - - -def apply_hf_chat_template( - tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast, - conversation: list[ConversationMessage], - chat_template: str | None, - tools: list[dict[str, Any]] | None, - *, - model_config: ModelConfig, - **kwargs: Any, -) -> str: - hf_chat_template = resolve_hf_chat_template( - tokenizer, - chat_template=chat_template, - tools=tools, - model_config=model_config, - ) - - if hf_chat_template is None: - raise ChatTemplateResolutionError( - "As of transformers v4.44, default chat template is no longer " - "allowed, so you must provide a chat template if the tokenizer " - "does not define one." - ) - - resolved_kwargs = resolve_chat_template_kwargs( - tokenizer=tokenizer, - chat_template=hf_chat_template, - chat_template_kwargs=kwargs, - ) - - try: - return tokenizer.apply_chat_template( - conversation=conversation, # type: ignore[arg-type] - tools=tools, # type: ignore[arg-type] - chat_template=hf_chat_template, - tokenize=False, - **resolved_kwargs, - ) - - # External library exceptions can sometimes occur despite the framework's - # internal exception management capabilities. - except Exception as e: - # Log and report any library-related exceptions for further - # investigation. - logger.exception( - "An error occurred in `transformers` while applying chat template" - ) - raise ValueError(str(e)) from e - - -def apply_mistral_chat_template( - tokenizer: "MistralTokenizer", - messages: list[ChatCompletionMessageParam], - chat_template: str | None, - tools: list[dict[str, Any]] | None, - **kwargs: Any, -) -> list[int]: - from mistral_common.exceptions import MistralCommonException - - # The return value of resolve_mistral_chat_template is always None, - # and we won't use it. - resolve_mistral_chat_template( - chat_template=chat_template, - **kwargs, - ) - - try: - return tokenizer.apply_chat_template( - messages=messages, - tools=tools, - **kwargs, - ) - # mistral-common uses assert statements to stop processing of input - # if input does not comply with the expected format. - # We convert those assertion errors to ValueErrors so they can be - # properly caught in the preprocessing_input step - except (AssertionError, MistralCommonException) as e: - raise ValueError(str(e)) from e - - # External library exceptions can sometimes occur despite the framework's - # internal exception management capabilities. - except Exception as e: - # Log and report any library-related exceptions for further - # investigation. - logger.exception( - "An error occurred in `mistral_common` while applying chat template" - ) - raise ValueError(str(e)) from e + return conversation, await mm_tracker.all_mm_data(), mm_tracker.all_mm_uuids() def get_history_tool_calls_cnt(conversation: list[ConversationMessage]): diff --git a/vllm/entrypoints/llm.py b/vllm/entrypoints/llm.py index e703ad5866c..ed7ad2c602f 100644 --- a/vllm/entrypoints/llm.py +++ b/vllm/entrypoints/llm.py @@ -37,10 +37,6 @@ from vllm.engine.arg_utils import EngineArgs from vllm.entrypoints.chat_utils import ( ChatCompletionMessageParam, ChatTemplateContentFormatOption, - apply_hf_chat_template, - apply_mistral_chat_template, - parse_chat_messages, - resolve_chat_template_content_format, ) from vllm.entrypoints.pooling.score.utils import ( ScoreContentPartParam, @@ -786,7 +782,7 @@ class LLM: tools: list[dict[str, Any]] | None = None, chat_template_kwargs: dict[str, Any] | None = None, mm_processor_kwargs: dict[str, Any] | None = None, - ) -> list[TokensPrompt]: + ) -> list[TextPrompt | TokensPrompt]: """ Generate prompt for a chat conversation. The pre-processed prompt can then be used as input for the other LLM methods. @@ -807,63 +803,27 @@ class LLM: # messages is list[...] list_of_messages = [cast(list[ChatCompletionMessageParam], messages)] - tokenizer = self.get_tokenizer() - model_config = self.model_config - resolved_content_format = resolve_chat_template_content_format( - chat_template, - tools, - chat_template_content_format, - tokenizer, - model_config=model_config, - ) + renderer = self.llm_engine.renderer - _chat_template_kwargs: dict[str, Any] = dict( - chat_template=chat_template, - add_generation_prompt=add_generation_prompt, - continue_final_message=continue_final_message, - tools=tools, - ) - _chat_template_kwargs.update(chat_template_kwargs or {}) + chat_template_kwargs = { + "chat_template": chat_template, + "add_generation_prompt": add_generation_prompt, + "continue_final_message": continue_final_message, + "tools": tools, + **(chat_template_kwargs or {}), + } - prompts: list[TokensPrompt] = [] + prompts = list[TextPrompt | TokensPrompt]() for msgs in list_of_messages: - # NOTE: _parse_chat_message_content_parts() currently doesn't + # NOTE: renderer.render_messages() currently doesn't # handle mm_processor_kwargs, since there is no implementation in # the chat message parsing for it. - conversation, mm_data, mm_uuids = parse_chat_messages( + _, prompt = renderer.render_messages( msgs, - model_config, - content_format=resolved_content_format, + chat_template_content_format=chat_template_content_format, + **chat_template_kwargs, ) - - if isinstance(tokenizer, MistralTokenizer): - prompt_token_ids = apply_mistral_chat_template( - tokenizer, - messages=msgs, - **_chat_template_kwargs, - ) - else: - prompt_str = apply_hf_chat_template( - tokenizer=tokenizer, - conversation=conversation, - model_config=model_config, - **_chat_template_kwargs, - ) - # Special tokens are already included in chat templates so - # should not be added by the tokenizer in this case. - prompt_token_ids = tokenizer.encode( - prompt_str, add_special_tokens=False - ) - - prompt = TokensPrompt(prompt_token_ids=prompt_token_ids) - - if mm_data is not None: - prompt["multi_modal_data"] = mm_data - - if mm_uuids is not None: - prompt["multi_modal_uuids"] = mm_uuids - if mm_processor_kwargs is not None: prompt["mm_processor_kwargs"] = mm_processor_kwargs diff --git a/vllm/entrypoints/openai/api_server.py b/vllm/entrypoints/openai/api_server.py index 9de6968ebb3..2d7a637df49 100644 --- a/vllm/entrypoints/openai/api_server.py +++ b/vllm/entrypoints/openai/api_server.py @@ -34,6 +34,7 @@ import vllm.envs as envs from vllm.engine.arg_utils import AsyncEngineArgs from vllm.engine.protocol import EngineClient from vllm.entrypoints.anthropic.serving import AnthropicServingMessages +from vllm.entrypoints.chat_utils import load_chat_template from vllm.entrypoints.launcher import serve_http from vllm.entrypoints.logger import RequestLogger from vllm.entrypoints.mcp.tool_server import DemoToolServer, MCPToolServer, ToolServer @@ -62,7 +63,6 @@ from vllm.entrypoints.serve.tokenize.serving import OpenAIServingTokenization from vllm.entrypoints.utils import ( cli_env_setup, log_non_default_args, - process_chat_template, process_lora_modules, sanitize_message, ) @@ -662,9 +662,7 @@ async def init_app_state( supported_tasks = await engine_client.get_supported_tasks() logger.info("Supported tasks: %s", supported_tasks) - resolved_chat_template = await process_chat_template( - args.chat_template, engine_client, vllm_config.model_config - ) + resolved_chat_template = load_chat_template(args.chat_template) if args.tool_server == "demo": tool_server: ToolServer | None = DemoToolServer() diff --git a/vllm/entrypoints/openai/chat_completion/serving.py b/vllm/entrypoints/openai/chat_completion/serving.py index ca2423855f9..a15c99c24c2 100644 --- a/vllm/entrypoints/openai/chat_completion/serving.py +++ b/vllm/entrypoints/openai/chat_completion/serving.py @@ -186,8 +186,7 @@ class OpenAIServingChat(OpenAIServing): start_time = time.perf_counter() try: - # Get the tokenizer from the engine - tokenizer = await self.engine_client.get_tokenizer() + renderer = self.engine_client.renderer # Create a minimal dummy request dummy_request = ChatCompletionRequest( @@ -203,7 +202,7 @@ class OpenAIServingChat(OpenAIServing): # 3. Tokenizer initialization for chat await self._preprocess_chat( dummy_request, - tokenizer, + renderer, dummy_request.messages, chat_template=self.chat_template, chat_template_content_format=self.chat_template_content_format, @@ -247,7 +246,8 @@ class OpenAIServingChat(OpenAIServing): raise self.engine_client.dead_error try: - tokenizer = await self.engine_client.get_tokenizer() + renderer = self.engine_client.renderer + tokenizer = renderer.tokenizer tool_parser = self.tool_parser @@ -308,7 +308,7 @@ class OpenAIServingChat(OpenAIServing): conversation, engine_prompts = await self._preprocess_chat( request, - tokenizer, + renderer, request.messages, chat_template=request.chat_template or self.chat_template, chat_template_content_format=self.chat_template_content_format, @@ -365,8 +365,6 @@ class OpenAIServingChat(OpenAIServing): ) model_name = self.models.model_name(lora_request) - - tokenizer = await self.engine_client.get_tokenizer() except (ValueError, TypeError, RuntimeError) as e: logger.exception("Error preparing request components") return self.create_error_response(e) @@ -463,6 +461,8 @@ class OpenAIServingChat(OpenAIServing): (result_generator,) = generators # Streaming response + tokenizer = self.renderer.tokenizer + if request.stream: return self.chat_completion_stream_generator( request, @@ -1784,7 +1784,7 @@ class OpenAIServingChat(OpenAIServing): else: if tokenizer is None: raise ValueError( - "Tokenizer not available when `skip_tokenizer_init=True`" + "Unable to get tokenizer because `skip_tokenizer_init=True`" ) token = tokenizer.decode(token_id) diff --git a/vllm/entrypoints/openai/completion/serving.py b/vllm/entrypoints/openai/completion/serving.py index c2bbedea985..fb14a2307ec 100644 --- a/vllm/entrypoints/openai/completion/serving.py +++ b/vllm/entrypoints/openai/completion/serving.py @@ -117,12 +117,7 @@ class OpenAIServingCompletion(OpenAIServing): ) try: - if self.model_config.skip_tokenizer_init: - tokenizer = None - else: - tokenizer = await self.engine_client.get_tokenizer() - renderer = self._get_renderer(tokenizer) - + renderer = self._get_completion_renderer() engine_prompts = await renderer.render_prompt_and_embeds( prompt_or_prompts=request.prompt, prompt_embeds=request.prompt_embeds, @@ -163,11 +158,6 @@ class OpenAIServingCompletion(OpenAIServing): try: lora_request = self._maybe_get_adapters(request) - - if self.model_config.skip_tokenizer_init: - tokenizer = None - else: - tokenizer = await self.engine_client.get_tokenizer() except (ValueError, TypeError, RuntimeError) as e: logger.exception("Error preparing request components") return self.create_error_response(e) @@ -280,6 +270,8 @@ class OpenAIServingCompletion(OpenAIServing): stream = request.stream and not request.use_beam_search # Streaming response + tokenizer = self.renderer.tokenizer + if stream: return self.completion_stream_generator( request, diff --git a/vllm/entrypoints/openai/engine/serving.py b/vllm/entrypoints/openai/engine/serving.py index 0f4ee51e799..2cf33328e6f 100644 --- a/vllm/entrypoints/openai/engine/serving.py +++ b/vllm/entrypoints/openai/engine/serving.py @@ -6,10 +6,9 @@ import sys import time import traceback from collections.abc import AsyncGenerator, Callable, Iterable, Mapping -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from http import HTTPStatus -from typing import Any, ClassVar, Generic, TypeAlias, TypeVar +from typing import Any, ClassVar, Generic, TypeAlias, TypeVar, cast import numpy as np from fastapi import Request @@ -26,10 +25,6 @@ from vllm.entrypoints.chat_utils import ( ChatCompletionMessageParam, ChatTemplateContentFormatOption, ConversationMessage, - apply_hf_chat_template, - apply_mistral_chat_template, - parse_chat_messages_futures, - resolve_chat_template_content_format, ) from vllm.entrypoints.logger import RequestLogger from vllm.entrypoints.openai.chat_completion.protocol import ( @@ -113,10 +108,9 @@ from vllm.multimodal import MultiModalDataDict from vllm.outputs import CompletionOutput, PoolingRequestOutput, RequestOutput from vllm.pooling_params import PoolingParams from vllm.reasoning import ReasoningParser, ReasoningParserManager +from vllm.renderers import RendererLike from vllm.sampling_params import BeamSearchParams, SamplingParams from vllm.tokenizers import TokenizerLike -from vllm.tokenizers.deepseek_v32 import DeepseekV32Tokenizer -from vllm.tokenizers.mistral import MistralTokenizer from vllm.tool_parsers import ToolParser, ToolParserManager from vllm.tracing import ( contains_trace_headers, @@ -127,10 +121,8 @@ from vllm.utils import random_uuid from vllm.utils.async_utils import ( AsyncMicrobatchTokenizer, collect_from_async_generator, - make_async, merge_async_iterators, ) -from vllm.utils.collection_utils import is_list_of from vllm.v1.engine import EngineCoreRequest @@ -215,7 +207,6 @@ class ResponseGenerationMixin: @dataclass(kw_only=True) class ServeContext(RequestProcessingMixin, ResponseGenerationMixin, Generic[RequestT]): - # Shared across all requests request: RequestT raw_request: Request | None = None model_name: str @@ -223,9 +214,6 @@ class ServeContext(RequestProcessingMixin, ResponseGenerationMixin, Generic[Requ created_time: int = field(default_factory=lambda: int(time.time())) lora_request: LoRARequest | None = None - # Shared across most requests - tokenizer: TokenizerLike | None = None - @dataclass(kw_only=True) class ClassificationServeContext(ServeContext[ClassificationRequest]): @@ -261,16 +249,13 @@ class OpenAIServing: self.request_logger = request_logger self.return_tokens_as_token_ids = return_tokens_as_token_ids - self._tokenizer_executor = ThreadPoolExecutor(max_workers=1) - self._apply_mistral_chat_template_async = make_async( - apply_mistral_chat_template, executor=self._tokenizer_executor - ) self._async_tokenizer_pool: dict[TokenizerLike, AsyncMicrobatchTokenizer] = {} self.log_error_stack = log_error_stack self.input_processor = self.models.input_processor self.io_processor = self.models.io_processor + self.renderer = self.models.renderer self.model_config = self.models.model_config self.max_model_len = self.model_config.max_model_len @@ -557,14 +542,14 @@ class OpenAIServing: prompt_logprobs=None, ) - def _get_renderer(self, tokenizer: TokenizerLike | None) -> BaseRenderer: + def _get_completion_renderer(self) -> BaseRenderer: """ Get a Renderer instance with the provided tokenizer. Uses shared async tokenizer pool for efficiency. """ return CompletionRenderer( model_config=self.model_config, - tokenizer=tokenizer, + tokenizer=self.renderer.tokenizer, async_tokenizer_pool=self._async_tokenizer_pool, ) @@ -1183,7 +1168,7 @@ class OpenAIServing: async def _preprocess_chat( self, request: ChatLikeRequest | ResponsesRequest, - tokenizer: TokenizerLike | None, + renderer: RendererLike, messages: list[ChatCompletionMessageParam], chat_template: str | None, chat_template_content_format: ChatTemplateContentFormatOption, @@ -1196,59 +1181,58 @@ class OpenAIServing: tool_parser: Callable[[TokenizerLike], ToolParser] | None = None, add_special_tokens: bool = False, ) -> tuple[list[ConversationMessage], list[TokensPrompt]]: - model_config = self.model_config - - resolved_content_format = resolve_chat_template_content_format( - chat_template, - tool_dicts, - chat_template_content_format, - tokenizer, - model_config=model_config, - ) - conversation, mm_data_future, mm_uuids = parse_chat_messages_futures( - messages, - model_config, - content_format=resolved_content_format, - ) - - _chat_template_kwargs: dict[str, Any] = dict( - chat_template=chat_template, - add_generation_prompt=add_generation_prompt, - continue_final_message=continue_final_message, - tools=tool_dicts, - documents=documents, - ) - _chat_template_kwargs |= self._prepare_extra_chat_template_kwargs( + chat_template_kwargs = { + "chat_template": chat_template, + "add_generation_prompt": add_generation_prompt, + "continue_final_message": continue_final_message, + "tools": tool_dicts, + "documents": documents, + **(chat_template_kwargs or {}), + } + chat_template_kwargs = self._prepare_extra_chat_template_kwargs( chat_template_kwargs, default_chat_template_kwargs, ) - request_prompt: str | list[int] + # Use the async tokenizer in `OpenAIServing` if possible. + # Later we can move it into the renderer so that we can return both + # text and token IDs in the same prompt from `render_messages_async` + # which is used for logging and `enable_response_messages`. + from vllm.tokenizers.mistral import MistralTokenizer - if tokenizer is None: - request_prompt = "placeholder" - elif isinstance(tokenizer, MistralTokenizer): - request_prompt = await self._apply_mistral_chat_template_async( - tokenizer, - messages=messages, - **_chat_template_kwargs, - ) - elif isinstance(tokenizer, DeepseekV32Tokenizer): - request_prompt = tokenizer.apply_chat_template( - conversation=conversation, - messages=messages, - model_config=model_config, - **_chat_template_kwargs, + conversation, engine_prompt = await renderer.render_messages_async( + messages, + chat_template_content_format=chat_template_content_format, + tokenize=( + chat_template_kwargs.pop("tokenize", False) + or isinstance(renderer.tokenizer, MistralTokenizer) + ), + **chat_template_kwargs, + ) + + if "prompt_token_ids" not in engine_prompt: + extra_data = engine_prompt + engine_prompt = await self._tokenize_prompt_input_async( + request, + renderer.get_tokenizer(), + engine_prompt["prompt"], + add_special_tokens=add_special_tokens, ) + # Fill in other keys like MM data + engine_prompt.update(extra_data) # type: ignore else: - request_prompt = apply_hf_chat_template( - tokenizer=tokenizer, - conversation=conversation, - model_config=model_config, - **_chat_template_kwargs, + self._validate_input( + request=request, + input_ids=engine_prompt["prompt_token_ids"], # type: ignore + input_text="", ) - mm_data = await mm_data_future + engine_prompt = cast(TokensPrompt, engine_prompt) + + if request.mm_processor_kwargs is not None: + engine_prompt["mm_processor_kwargs"] = request.mm_processor_kwargs + if (cache_salt := getattr(request, "cache_salt", None)) is not None: + engine_prompt["cache_salt"] = cache_salt # tool parsing is done only if a tool_parser has been set and if # tool_choice is not "none" (if tool_choice is "none" but a tool_parser @@ -1264,49 +1248,10 @@ class OpenAIServing: "or Responses API requests." ) raise NotImplementedError(msg) + + tokenizer = renderer.get_tokenizer() request = tool_parser(tokenizer).adjust_request(request=request) # type: ignore - if tokenizer is None: - assert isinstance(request_prompt, str), ( - "Prompt has to be a string", - "when the tokenizer is not initialised", - ) - prompt_inputs = TokensPrompt(prompt=request_prompt, prompt_token_ids=[1]) - elif isinstance(request_prompt, str): - prompt_inputs = await self._tokenize_prompt_input_async( - request, - tokenizer, - request_prompt, - add_special_tokens=add_special_tokens, - ) - else: - # For MistralTokenizer - assert is_list_of(request_prompt, int), ( - "Prompt has to be either a string or a list of token ids" - ) - input_text = tokenizer.decode(request_prompt) - prompt_inputs = self._validate_input( - request=request, - input_ids=request_prompt, - input_text=input_text, - ) - - engine_prompt = TokensPrompt(prompt_token_ids=prompt_inputs["prompt_token_ids"]) - if "prompt" in prompt_inputs: - engine_prompt["prompt"] = prompt_inputs["prompt"] - - if mm_data is not None: - engine_prompt["multi_modal_data"] = mm_data - - if mm_uuids is not None: - engine_prompt["multi_modal_uuids"] = mm_uuids - - if request.mm_processor_kwargs is not None: - engine_prompt["mm_processor_kwargs"] = request.mm_processor_kwargs - - if hasattr(request, "cache_salt") and request.cache_salt is not None: - engine_prompt["cache_salt"] = request.cache_salt - return conversation, [engine_prompt] async def _process_inputs( @@ -1341,7 +1286,7 @@ class OpenAIServing: async def _render_next_turn( self, request: ResponsesRequest, - tokenizer: TokenizerLike | None, + renderer: RendererLike, messages: list[ResponseInputOutputItem], tool_dicts: list[dict[str, Any]] | None, tool_parser, @@ -1354,7 +1299,7 @@ class OpenAIServing: _, engine_prompts = await self._preprocess_chat( request, - tokenizer, + renderer, new_messages, tool_dicts=tool_dicts, tool_parser=tool_parser, @@ -1431,7 +1376,7 @@ class OpenAIServing: elif isinstance(context, ParsableContext): engine_prompts = await self._render_next_turn( context.request, - context.tokenizer, + context.renderer, context.parser.response_messages, context.tool_dicts, context.tool_parser_cls, diff --git a/vllm/entrypoints/openai/models/serving.py b/vllm/entrypoints/openai/models/serving.py index a4b92e5ec33..ba32787acd6 100644 --- a/vllm/entrypoints/openai/models/serving.py +++ b/vllm/entrypoints/openai/models/serving.py @@ -61,6 +61,7 @@ class OpenAIServingModels: self.input_processor = self.engine_client.input_processor self.io_processor = self.engine_client.io_processor + self.renderer = self.engine_client.renderer self.model_config = self.engine_client.model_config self.max_model_len = self.model_config.max_model_len diff --git a/vllm/entrypoints/openai/responses/context.py b/vllm/entrypoints/openai/responses/context.py index d6c818aabbe..b3ac24881aa 100644 --- a/vllm/entrypoints/openai/responses/context.py +++ b/vllm/entrypoints/openai/responses/context.py @@ -43,6 +43,7 @@ from vllm.entrypoints.openai.responses.protocol import ( from vllm.entrypoints.openai.responses.utils import construct_tool_dicts from vllm.outputs import RequestOutput from vllm.reasoning.abs_reasoning_parsers import ReasoningParser +from vllm.renderers import RendererLike from vllm.tokenizers import TokenizerLike from vllm.tool_parsers.abstract_tool_parser import ToolParser from vllm.utils import random_uuid @@ -260,7 +261,7 @@ class ParsableContext(ConversationContext): self, *, response_messages: list[ResponseInputOutputItem], - tokenizer: TokenizerLike, + renderer: RendererLike, reasoning_parser_cls: Callable[[TokenizerLike], ReasoningParser] | None, request: ResponsesRequest, available_tools: list[str] | None, @@ -279,6 +280,7 @@ class ParsableContext(ConversationContext): if reasoning_parser_cls is None: raise ValueError("reasoning_parser_cls must be provided.") + tokenizer = renderer.get_tokenizer() self.parser = get_responses_parser_for_simple_context( tokenizer=tokenizer, reasoning_parser_cls=reasoning_parser_cls, @@ -288,6 +290,7 @@ class ParsableContext(ConversationContext): ) self.tool_parser_cls = tool_parser_cls self.request = request + self.renderer = renderer self.tokenizer = tokenizer self.available_tools = available_tools or [] diff --git a/vllm/entrypoints/openai/responses/serving.py b/vllm/entrypoints/openai/responses/serving.py index 1a74e6ae063..cb0317f9f85 100644 --- a/vllm/entrypoints/openai/responses/serving.py +++ b/vllm/entrypoints/openai/responses/serving.py @@ -121,6 +121,7 @@ from vllm.logger import init_logger from vllm.logprobs import Logprob as SampleLogprob from vllm.logprobs import SampleLogprobs from vllm.outputs import CompletionOutput +from vllm.renderers import RendererLike from vllm.sampling_params import SamplingParams, StructuredOutputsParams from vllm.tokenizers import TokenizerLike from vllm.utils import random_uuid @@ -380,7 +381,8 @@ class OpenAIServingResponses(OpenAIServing): try: lora_request = self._maybe_get_adapters(request) model_name = self.models.model_name(lora_request) - tokenizer = await self.engine_client.get_tokenizer() + renderer = self.engine_client.renderer + tokenizer = renderer.get_tokenizer() if self.use_harmony: messages, engine_prompts = self._make_request_with_harmony( @@ -388,7 +390,7 @@ class OpenAIServingResponses(OpenAIServing): ) else: messages, engine_prompts = await self._make_request( - request, prev_response, tokenizer + request, prev_response, renderer ) except ( @@ -454,7 +456,7 @@ class OpenAIServingResponses(OpenAIServing): # tokens during generation instead of at the end context = ParsableContext( response_messages=messages, - tokenizer=tokenizer, + renderer=renderer, reasoning_parser_cls=self.reasoning_parser, request=request, tool_parser_cls=self.tool_parser, @@ -585,7 +587,7 @@ class OpenAIServingResponses(OpenAIServing): self, request: ResponsesRequest, prev_response: ResponsesResponse | None, - tokenizer: TokenizerLike, + renderer: RendererLike, ): tool_dicts = construct_tool_dicts(request.tools, request.tool_choice) # Construct the input messages. @@ -607,7 +609,7 @@ class OpenAIServingResponses(OpenAIServing): _, engine_prompts = await self._preprocess_chat( request, - tokenizer, + renderer, messages, tool_dicts=tool_dicts, tool_parser=self.tool_parser, @@ -631,6 +633,7 @@ class OpenAIServingResponses(OpenAIServing): raise NotImplementedError( "Only 'auto' tool_choice is supported in response API with Harmony" ) + messages = self._construct_input_messages_with_harmony(request, prev_response) prompt_token_ids = render_for_completion(messages) engine_prompt = TokensPrompt(prompt_token_ids=prompt_token_ids) diff --git a/vllm/entrypoints/pooling/__init__.py b/vllm/entrypoints/pooling/__init__.py index e9b2139b15b..408542dfa52 100644 --- a/vllm/entrypoints/pooling/__init__.py +++ b/vllm/entrypoints/pooling/__init__.py @@ -28,21 +28,17 @@ def register_pooling_api_routers(app: FastAPI): async def init_pooling_state( engine_client: "EngineClient", state: "State", args: "Namespace" ): + from vllm.entrypoints.chat_utils import load_chat_template from vllm.entrypoints.logger import RequestLogger from vllm.entrypoints.pooling.classify.serving import ServingClassification from vllm.entrypoints.pooling.embed.serving import OpenAIServingEmbedding from vllm.entrypoints.pooling.pooling.serving import OpenAIServingPooling from vllm.entrypoints.pooling.score.serving import ServingScores - from vllm.entrypoints.utils import process_chat_template from vllm.tasks import POOLING_TASKS supported_tasks = await engine_client.get_supported_tasks() - vllm_config = engine_client.vllm_config - - resolved_chat_template = await process_chat_template( - args.chat_template, engine_client, vllm_config.model_config - ) + resolved_chat_template = load_chat_template(args.chat_template) if args.enable_log_requests: request_logger = RequestLogger(max_log_len=args.max_log_len) diff --git a/vllm/entrypoints/pooling/classify/serving.py b/vllm/entrypoints/pooling/classify/serving.py index 2ff3139302f..cd8aa18e854 100644 --- a/vllm/entrypoints/pooling/classify/serving.py +++ b/vllm/entrypoints/pooling/classify/serving.py @@ -54,8 +54,6 @@ class ClassificationMixin(OpenAIServing): """ ctx = cast(ClassificationServeContext, ctx) try: - ctx.tokenizer = await self.engine_client.get_tokenizer() - request_obj = ctx.request if isinstance(request_obj, ClassificationChatRequest): @@ -76,7 +74,7 @@ class ClassificationMixin(OpenAIServing): _, engine_prompts = await self._preprocess_chat( cast(ChatCompletionRequest, chat_request), - ctx.tokenizer, + self.renderer, messages, chat_template=( chat_request.chat_template @@ -104,7 +102,7 @@ class ClassificationMixin(OpenAIServing): ctx.engine_prompts = [] return None - renderer = self._get_renderer(ctx.tokenizer) + renderer = self._get_completion_renderer() prompt_input = cast(str | list[str], input_data) ctx.engine_prompts = await renderer.render_prompt( prompt_or_prompts=prompt_input, diff --git a/vllm/entrypoints/pooling/embed/serving.py b/vllm/entrypoints/pooling/embed/serving.py index b48e3a016ae..4251276acf7 100644 --- a/vllm/entrypoints/pooling/embed/serving.py +++ b/vllm/entrypoints/pooling/embed/serving.py @@ -78,13 +78,10 @@ class EmbeddingMixin(OpenAIServing): try: ctx.lora_request = self._maybe_get_adapters(ctx.request) - tokenizer = await self.engine_client.get_tokenizer() - renderer = self._get_renderer(tokenizer) - if isinstance(ctx.request, EmbeddingChatRequest): _, ctx.engine_prompts = await self._preprocess_chat( ctx.request, - tokenizer, + self.renderer, ctx.request.messages, chat_template=ctx.request.chat_template or ctx.chat_template, chat_template_content_format=ctx.chat_template_content_format, @@ -93,6 +90,7 @@ class EmbeddingMixin(OpenAIServing): add_special_tokens=ctx.request.add_special_tokens, ) else: + renderer = self._get_completion_renderer() ctx.engine_prompts = await renderer.render_prompt( prompt_or_prompts=ctx.request.input, config=self._build_render_config(ctx.request), diff --git a/vllm/entrypoints/pooling/pooling/serving.py b/vllm/entrypoints/pooling/pooling/serving.py index b53caf81b27..1900e446dbb 100644 --- a/vllm/entrypoints/pooling/pooling/serving.py +++ b/vllm/entrypoints/pooling/pooling/serving.py @@ -94,12 +94,6 @@ class OpenAIServingPooling(OpenAIServing): try: lora_request = self._maybe_get_adapters(request) - if self.model_config.skip_tokenizer_init: - tokenizer = None - else: - tokenizer = await self.engine_client.get_tokenizer() - renderer = self._get_renderer(tokenizer) - if getattr(request, "dimensions", None) is not None: return self.create_error_response( "dimensions is currently not supported" @@ -140,7 +134,7 @@ class OpenAIServingPooling(OpenAIServing): _, engine_prompts = await self._preprocess_chat( request, - tokenizer, + self.renderer, request.messages, chat_template=request.chat_template or self.chat_template, chat_template_content_format=self.chat_template_content_format, @@ -149,6 +143,7 @@ class OpenAIServingPooling(OpenAIServing): add_special_tokens=request.add_special_tokens, ) elif isinstance(request, PoolingCompletionRequest): + renderer = self._get_completion_renderer() engine_prompts = await renderer.render_prompt( prompt_or_prompts=request.input, config=self._build_render_config(request), diff --git a/vllm/entrypoints/pooling/score/serving.py b/vllm/entrypoints/pooling/score/serving.py index 1040d2be107..85c74e5a26c 100644 --- a/vllm/entrypoints/pooling/score/serving.py +++ b/vllm/entrypoints/pooling/score/serving.py @@ -3,6 +3,7 @@ import asyncio import time from collections.abc import AsyncGenerator, Mapping +from concurrent.futures import ThreadPoolExecutor from typing import Any from fastapi import Request @@ -63,6 +64,8 @@ class ServingScores(OpenAIServing): ) self.score_template = score_template + self._tokenizer_executor = ThreadPoolExecutor(max_workers=1) + async def _embedding_score( self, tokenizer: TokenizerLike, @@ -283,8 +286,7 @@ class ServingScores(OpenAIServing): raw_request: Request | None = None, ) -> list[PoolingRequestOutput] | ErrorResponse: lora_request = self._maybe_get_adapters(request) - - tokenizer = await self.engine_client.get_tokenizer() + tokenizer = self.renderer.get_tokenizer() truncate_prompt_tokens = getattr(request, "truncate_prompt_tokens", None) diff --git a/vllm/entrypoints/pooling/score/utils.py b/vllm/entrypoints/pooling/score/utils.py index 09ef8781b2d..8fac0dd8b39 100644 --- a/vllm/entrypoints/pooling/score/utils.py +++ b/vllm/entrypoints/pooling/score/utils.py @@ -16,12 +16,12 @@ from vllm.entrypoints.chat_utils import ( MultiModalItemTracker, _ContentPart, _parse_chat_message_content_part, - apply_hf_chat_template, ) from vllm.inputs import TokensPrompt from vllm.model_executor.models.interfaces import supports_score_template from vllm.multimodal.inputs import MultiModalDataDict from vllm.outputs import PoolingRequestOutput +from vllm.renderers.hf import safe_apply_chat_template from vllm.tokenizers import TokenizerLike ScoreContentPartParam: TypeAlias = ( @@ -224,15 +224,16 @@ def get_score_prompt( # If that fails because there is no such template, # fall back to the default implementation. try: - full_prompt = apply_hf_chat_template( + full_prompt = safe_apply_chat_template( + model_config, tokenizer, [ {"role": "query", "content": prompt_1}, {"role": "document", "content": prompt_2}, ], - score_template, + chat_template=score_template, tools=None, - model_config=model_config, + tokenize=False, ) prompt_inputs = tokenizer(full_prompt, **tokenization_kwargs) except ChatTemplateResolutionError: diff --git a/vllm/entrypoints/serve/tokenize/serving.py b/vllm/entrypoints/serve/tokenize/serving.py index b57c18bf56c..f0cfb8af174 100644 --- a/vllm/entrypoints/serve/tokenize/serving.py +++ b/vllm/entrypoints/serve/tokenize/serving.py @@ -67,9 +67,6 @@ class OpenAIServingTokenization(OpenAIServing): try: lora_request = self._maybe_get_adapters(request) - tokenizer = await self.engine_client.get_tokenizer() - renderer = self._get_renderer(tokenizer) - if isinstance(request, TokenizeChatRequest): tool_dicts = ( None @@ -86,7 +83,7 @@ class OpenAIServingTokenization(OpenAIServing): _, engine_prompts = await self._preprocess_chat( request, - tokenizer, + self.renderer, request.messages, tool_dicts=tool_dicts, chat_template=request.chat_template or self.chat_template, @@ -97,6 +94,7 @@ class OpenAIServingTokenization(OpenAIServing): add_special_tokens=request.add_special_tokens, ) else: + renderer = self._get_completion_renderer() engine_prompts = await renderer.render_prompt( prompt_or_prompts=request.prompt, config=self._build_render_config(request), @@ -116,6 +114,7 @@ class OpenAIServingTokenization(OpenAIServing): token_strs = None if request.return_token_strs: + tokenizer = self.renderer.get_tokenizer() token_strs = tokenizer.convert_ids_to_tokens(input_ids) return TokenizeResponse( @@ -137,8 +136,7 @@ class OpenAIServingTokenization(OpenAIServing): request_id = f"tokenize-{self._base_request_id(raw_request)}" lora_request = self._maybe_get_adapters(request) - - tokenizer = await self.engine_client.get_tokenizer() + tokenizer = self.renderer.get_tokenizer() self._log_inputs( request_id, @@ -161,7 +159,7 @@ class OpenAIServingTokenization(OpenAIServing): ) -> TokenizerInfoResponse | ErrorResponse: """Get comprehensive tokenizer information.""" try: - tokenizer = await self.engine_client.get_tokenizer() + tokenizer = self.renderer.get_tokenizer() info = TokenizerInfo(tokenizer, self.chat_template).to_dict() return TokenizerInfoResponse(**info) except Exception as e: diff --git a/vllm/entrypoints/utils.py b/vllm/entrypoints/utils.py index 9fb21484fa1..64cbb6ed442 100644 --- a/vllm/entrypoints/utils.py +++ b/vllm/entrypoints/utils.py @@ -6,7 +6,6 @@ import dataclasses import functools import os from argparse import Namespace -from pathlib import Path from typing import TYPE_CHECKING, Any import regex as re @@ -14,17 +13,9 @@ from fastapi import Request from fastapi.responses import JSONResponse, StreamingResponse from starlette.background import BackgroundTask, BackgroundTasks -from vllm.config import ModelConfig from vllm.engine.arg_utils import EngineArgs -from vllm.engine.protocol import EngineClient -from vllm.entrypoints.chat_utils import ( - load_chat_template, - resolve_hf_chat_template, - resolve_mistral_chat_template, -) from vllm.logger import init_logger from vllm.platforms import current_platform -from vllm.tokenizers.mistral import MistralTokenizer from vllm.utils.argparse_utils import FlexibleArgumentParser if TYPE_CHECKING: @@ -301,40 +292,6 @@ def process_lora_modules( return lora_modules -async def process_chat_template( - args_chat_template: Path | str | None, - engine_client: EngineClient, - model_config: ModelConfig, -) -> str | None: - resolved_chat_template = load_chat_template(args_chat_template) - if resolved_chat_template is not None: - # Get the tokenizer to check official template - tokenizer = await engine_client.get_tokenizer() - - if isinstance(tokenizer, MistralTokenizer): - # The warning is logged in resolve_mistral_chat_template. - resolved_chat_template = resolve_mistral_chat_template( - chat_template=resolved_chat_template - ) - else: - hf_chat_template = resolve_hf_chat_template( - tokenizer=tokenizer, - chat_template=None, - tools=None, - model_config=model_config, - ) - - if hf_chat_template != resolved_chat_template: - logger.warning( - "Using supplied chat template: %s\n" - "It is different from official chat template '%s'. " - "This discrepancy may lead to performance degradation.", - resolved_chat_template, - model_config.model, - ) - return resolved_chat_template - - def sanitize_message(message: str) -> str: # Avoid leaking memory address from object reprs return re.sub(r" at 0x[0-9a-f]+>", ">", message) diff --git a/vllm/inputs/preprocess.py b/vllm/inputs/preprocess.py index 6723809b51e..eb0a38f51ea 100644 --- a/vllm/inputs/preprocess.py +++ b/vllm/inputs/preprocess.py @@ -17,6 +17,7 @@ from vllm.multimodal.inputs import ( MultiModalUUIDDict, ) from vllm.multimodal.processing import BaseMultiModalProcessor +from vllm.renderers import renderer_from_config from vllm.tokenizers import TokenizerLike from vllm.utils.jsontree import json_iter_leaves from vllm.v1.metrics.stats import MultiModalCacheStats @@ -46,7 +47,6 @@ class InputPreprocessor: def __init__( self, model_config: ModelConfig, - tokenizer: TokenizerLike | None, observability_config: ObservabilityConfig | None = None, mm_registry: MultiModalRegistry = MULTIMODAL_REGISTRY, mm_processor_cache: BaseMultiModalProcessorCache | None = None, @@ -54,20 +54,19 @@ class InputPreprocessor: super().__init__() self.model_config = model_config - self.tokenizer = tokenizer self.observability_config = observability_config + self.renderer = renderer_from_config(model_config) self.mm_registry = mm_registry self.mm_processor_cache = mm_processor_cache self.mm_cache_stats = MultiModalCacheStats() if mm_processor_cache else None - def get_tokenizer(self) -> TokenizerLike: - if self.tokenizer is None: - raise ValueError( - "You cannot pass text prompts when `skip_tokenizer_init=True`" - ) + @property + def tokenizer(self) -> TokenizerLike | None: + return self.renderer.tokenizer - return self.tokenizer + def get_tokenizer(self) -> TokenizerLike: + return self.renderer.get_tokenizer() def get_bos_token_id(self) -> int | None: if self.tokenizer is None: diff --git a/vllm/renderers/__init__.py b/vllm/renderers/__init__.py new file mode 100644 index 00000000000..cd6a11dcc83 --- /dev/null +++ b/vllm/renderers/__init__.py @@ -0,0 +1,7 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from .protocol import RendererLike +from .registry import RendererRegistry, renderer_from_config + +__all__ = ["RendererLike", "RendererRegistry", "renderer_from_config"] diff --git a/vllm/renderers/deepseek_v32.py b/vllm/renderers/deepseek_v32.py new file mode 100644 index 00000000000..123911654d8 --- /dev/null +++ b/vllm/renderers/deepseek_v32.py @@ -0,0 +1,119 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ConversationMessage, + parse_chat_messages, + parse_chat_messages_async, +) +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.logger import init_logger +from vllm.tokenizers import cached_get_tokenizer +from vllm.tokenizers.deepseek_v32 import DeepseekV32Tokenizer + +from .protocol import RendererLike + +logger = init_logger(__name__) + + +class DeepseekV32Renderer(RendererLike): + @classmethod + def from_config( + cls, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + return cls(config, tokenizer_kwargs) + + def __init__( + self, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> None: + super().__init__() + + self.config = config + + if config.skip_tokenizer_init: + tokenizer = None + else: + tokenizer = cached_get_tokenizer( + tokenizer_cls=DeepseekV32Tokenizer, + **tokenizer_kwargs, + ) + + self._tokenizer = tokenizer + + @property + def tokenizer(self) -> DeepseekV32Tokenizer | None: + return self._tokenizer + + def get_tokenizer(self) -> DeepseekV32Tokenizer: + tokenizer = self.tokenizer + if tokenizer is None: + raise ValueError("Tokenizer not available when `skip_tokenizer_init=True`") + + return tokenizer + + def render_messages( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = parse_chat_messages( + messages, + self.config, + content_format="string", + ) + + prompt_raw = tokenizer.apply_chat_template( + conversation=conversation, + messages=messages, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] + + async def render_messages_async( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = await parse_chat_messages_async( + messages, + self.config, + content_format="string", + ) + + prompt_raw = tokenizer.apply_chat_template( + conversation=conversation, + messages=messages, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] diff --git a/vllm/renderers/grok2.py b/vllm/renderers/grok2.py new file mode 100644 index 00000000000..06de760f8f9 --- /dev/null +++ b/vllm/renderers/grok2.py @@ -0,0 +1,119 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ConversationMessage, + parse_chat_messages, + parse_chat_messages_async, +) +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.logger import init_logger +from vllm.tokenizers import cached_get_tokenizer +from vllm.tokenizers.grok2 import Grok2Tokenizer + +from .protocol import RendererLike + +logger = init_logger(__name__) + + +class Grok2Renderer(RendererLike): + @classmethod + def from_config( + cls, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + return cls(config, tokenizer_kwargs) + + def __init__( + self, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> None: + super().__init__() + + self.config = config + + if config.skip_tokenizer_init: + tokenizer = None + else: + tokenizer = cached_get_tokenizer( + tokenizer_cls=Grok2Tokenizer, + **tokenizer_kwargs, + ) + + self._tokenizer = tokenizer + + @property + def tokenizer(self) -> Grok2Tokenizer | None: + return self._tokenizer + + def get_tokenizer(self) -> Grok2Tokenizer: + tokenizer = self.tokenizer + if tokenizer is None: + raise ValueError("Tokenizer not available when `skip_tokenizer_init=True`") + + return tokenizer + + def render_messages( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = parse_chat_messages( + messages, + self.config, + content_format="string", + ) + + prompt_raw = tokenizer.apply_chat_template( + conversation=conversation, + messages=messages, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] + + async def render_messages_async( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = await parse_chat_messages_async( + messages, + self.config, + content_format="string", + ) + + prompt_raw = tokenizer.apply_chat_template( + conversation=conversation, + messages=messages, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] diff --git a/vllm/renderers/hf.py b/vllm/renderers/hf.py new file mode 100644 index 00000000000..d2252c65544 --- /dev/null +++ b/vllm/renderers/hf.py @@ -0,0 +1,600 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import inspect +from collections import deque +from collections.abc import Set +from functools import lru_cache +from typing import Any, cast + +import jinja2 +import jinja2.ext +import jinja2.meta +import jinja2.nodes +import jinja2.parser +import jinja2.sandbox + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ChatTemplateContentFormat, + ChatTemplateContentFormatOption, + ChatTemplateResolutionError, + ConversationMessage, + load_chat_template, + parse_chat_messages, + parse_chat_messages_async, +) +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.logger import init_logger +from vllm.tokenizers import cached_get_tokenizer +from vllm.tokenizers.hf import CachedHfTokenizer, HfTokenizer +from vllm.transformers_utils.chat_templates import get_chat_template_fallback_path +from vllm.transformers_utils.processor import cached_get_processor +from vllm.utils.func_utils import supports_kw + +from .protocol import RendererLike + +logger = init_logger(__name__) + + +_PROCESSOR_CHAT_TEMPLATES = dict[tuple[str, bool], str | None]() +""" +Used in `_try_get_processor_chat_template` to avoid calling +`cached_get_processor` again if the processor fails to be loaded. + +This is needed because `lru_cache` does not cache when an exception happens. +""" + + +def _try_get_processor_chat_template( + tokenizer: HfTokenizer, + *, + trust_remote_code: bool, +) -> str | None: + cache_key = (tokenizer.name_or_path, trust_remote_code) + if cache_key in _PROCESSOR_CHAT_TEMPLATES: + return _PROCESSOR_CHAT_TEMPLATES[cache_key] + + from transformers import ( + PreTrainedTokenizer, + PreTrainedTokenizerFast, + ProcessorMixin, + ) + + try: + processor = cached_get_processor( + tokenizer.name_or_path, + processor_cls=( + PreTrainedTokenizer, + PreTrainedTokenizerFast, + ProcessorMixin, + ), + trust_remote_code=trust_remote_code, + ) + if ( + isinstance(processor, ProcessorMixin) + and hasattr(processor, "chat_template") + and (chat_template := processor.chat_template) is not None + ): + _PROCESSOR_CHAT_TEMPLATES[cache_key] = chat_template + return chat_template + except Exception: + logger.debug( + "Failed to load AutoProcessor chat template for %s", + tokenizer.name_or_path, + exc_info=True, + ) + + _PROCESSOR_CHAT_TEMPLATES[cache_key] = None + return None + + +def resolve_chat_template( + tokenizer: HfTokenizer, + chat_template: str | None, + tools: list[dict[str, Any]] | None, + *, + model_config: "ModelConfig", +) -> str | None: + # 1st priority: The given chat template + if chat_template is not None: + return chat_template + + # 2nd priority: AutoProcessor chat template, unless tool calling is enabled + if tools is None: + chat_template = _try_get_processor_chat_template( + tokenizer, + trust_remote_code=model_config.trust_remote_code, + ) + if chat_template is not None: + return chat_template + + # 3rd priority: AutoTokenizer chat template + try: + return tokenizer.get_chat_template(chat_template, tools=tools) + except Exception: + logger.debug( + "Failed to load AutoTokenizer chat template for %s", + tokenizer.name_or_path, + exc_info=True, + ) + + # 4th priority: Predefined fallbacks + path = get_chat_template_fallback_path( + model_type=model_config.hf_config.model_type, + tokenizer_name_or_path=tokenizer.name_or_path, + ) + if path is not None: + logger.info_once( + "Loading chat template fallback for %s as there isn't one " + "defined on HF Hub.", + tokenizer.name_or_path, + ) + chat_template = load_chat_template(path) + else: + logger.debug_once( + "There is no chat template fallback for %s", tokenizer.name_or_path + ) + + return chat_template + + +def _is_var_access(node: jinja2.nodes.Node, varname: str) -> bool: + if isinstance(node, jinja2.nodes.Name): + return node.ctx == "load" and node.name == varname + + return False + + +def _is_attr_access(node: jinja2.nodes.Node, varname: str, key: str) -> bool: + if isinstance(node, jinja2.nodes.Getitem): + return ( + _is_var_access(node.node, varname) + and isinstance(node.arg, jinja2.nodes.Const) + and node.arg.value == key + ) + + if isinstance(node, jinja2.nodes.Getattr): + return _is_var_access(node.node, varname) and node.attr == key + + return False + + +def _is_var_or_elems_access( + node: jinja2.nodes.Node, + varname: str, + key: str | None = None, +) -> bool: + if isinstance(node, jinja2.nodes.Filter): + return node.node is not None and _is_var_or_elems_access( + node.node, varname, key + ) + if isinstance(node, jinja2.nodes.Test): + return _is_var_or_elems_access(node.node, varname, key) + + if isinstance(node, jinja2.nodes.Getitem) and isinstance( + node.arg, jinja2.nodes.Slice + ): + return _is_var_or_elems_access(node.node, varname, key) + + return _is_attr_access(node, varname, key) if key else _is_var_access(node, varname) + + +def _iter_nodes_assign_var_or_elems(root: jinja2.nodes.Node, varname: str): + # Global variable that is implicitly defined at the root + yield root, varname + + # Iterative BFS + related_varnames = deque([varname]) + while related_varnames: + related_varname = related_varnames.popleft() + + for assign_ast in root.find_all(jinja2.nodes.Assign): + lhs = assign_ast.target + rhs = assign_ast.node + + if _is_var_or_elems_access(rhs, related_varname): + assert isinstance(lhs, jinja2.nodes.Name) + yield assign_ast, lhs.name + + # Avoid infinite looping for self-assignment + if lhs.name != related_varname: + related_varnames.append(lhs.name) + + +# NOTE: The proper way to handle this is to build a CFG so that we can handle +# the scope in which each variable is defined, but that is too complicated +def _iter_nodes_assign_messages_item(root: jinja2.nodes.Node): + messages_varnames = [ + varname for _, varname in _iter_nodes_assign_var_or_elems(root, "messages") + ] + + # Search for {%- for message in messages -%} loops + for loop_ast in root.find_all(jinja2.nodes.For): + loop_iter = loop_ast.iter + loop_target = loop_ast.target + + for varname in messages_varnames: + if _is_var_or_elems_access(loop_iter, varname): + assert isinstance(loop_target, jinja2.nodes.Name) + yield loop_ast, loop_target.name + break + + +def _iter_nodes_assign_content_item(root: jinja2.nodes.Node): + message_varnames = [ + varname for _, varname in _iter_nodes_assign_messages_item(root) + ] + + # Search for {%- for content in message['content'] -%} loops + for loop_ast in root.find_all(jinja2.nodes.For): + loop_iter = loop_ast.iter + loop_target = loop_ast.target + + for varname in message_varnames: + if _is_var_or_elems_access(loop_iter, varname, "content"): + assert isinstance(loop_target, jinja2.nodes.Name) + yield loop_ast, loop_target.name + break + + +def _try_extract_ast(chat_template: str) -> jinja2.nodes.Template | None: + import transformers.utils.chat_template_utils as hf_chat_utils + + try: + jinja_compiled = hf_chat_utils._compile_jinja_template(chat_template) + return jinja_compiled.environment.parse(chat_template) + except Exception: + logger.exception("Error when compiling Jinja template") + return None + + +@lru_cache(maxsize=32) +def _detect_content_format( + chat_template: str, + *, + default: ChatTemplateContentFormat, +) -> ChatTemplateContentFormat: + jinja_ast = _try_extract_ast(chat_template) + if jinja_ast is None: + return default + + try: + next(_iter_nodes_assign_content_item(jinja_ast)) + except StopIteration: + return "string" + except Exception: + logger.exception("Error when parsing AST of Jinja template") + return default + else: + return "openai" + + +def _resolve_chat_template_content_format( + chat_template: str | None, + tools: list[dict[str, Any]] | None, + tokenizer: HfTokenizer, + *, + model_config: "ModelConfig", +) -> ChatTemplateContentFormat: + resolved_chat_template = resolve_chat_template( + tokenizer, + chat_template=chat_template, + tools=tools, + model_config=model_config, + ) + + jinja_text = ( + resolved_chat_template + if isinstance(resolved_chat_template, str) + else load_chat_template(chat_template, is_literal=True) + ) + + detected_format = ( + "string" + if jinja_text is None + else _detect_content_format(jinja_text, default="string") + ) + + return detected_format + + +@lru_cache +def _log_chat_template_content_format( + chat_template: str | None, # For caching purposes + given_format: ChatTemplateContentFormatOption, + detected_format: ChatTemplateContentFormatOption, +): + logger.info( + "Detected the chat template content format to be '%s'. " + "You can set `--chat-template-content-format` to override this.", + detected_format, + ) + + if given_format != "auto" and given_format != detected_format: + logger.warning( + "You specified `--chat-template-content-format %s` " + "which is different from the detected format '%s'. " + "If our automatic detection is incorrect, please consider " + "opening a GitHub issue so that we can improve it: " + "https://github.com/vllm-project/vllm/issues/new/choose", + given_format, + detected_format, + ) + + +def resolve_chat_template_content_format( + chat_template: str | None, + tools: list[dict[str, Any]] | None, + given_format: ChatTemplateContentFormatOption, + tokenizer: HfTokenizer, + *, + model_config: "ModelConfig", +) -> ChatTemplateContentFormat: + if given_format != "auto": + return given_format + + detected_format = _resolve_chat_template_content_format( + chat_template, + tools, + tokenizer, + model_config=model_config, + ) + + _log_chat_template_content_format( + chat_template, + given_format=given_format, + detected_format=detected_format, + ) + + return detected_format + + +# adapted from https://github.com/huggingface/transformers/blob/v4.56.2/src/transformers/utils/chat_template_utils.py#L398-L412 +# only preserve the parse function used to resolve chat template kwargs +class AssistantTracker(jinja2.ext.Extension): + tags = {"generation"} + + def parse(self, parser: jinja2.parser.Parser) -> jinja2.nodes.Node: + lineno = next(parser.stream).lineno + body = parser.parse_statements(("name:endgeneration",), drop_needle=True) + call = self.call_method("_generation_support") + call_block = jinja2.nodes.CallBlock(call, [], [], body) + return call_block.set_lineno(lineno) + + +def _resolve_chat_template_kwargs(chat_template: str) -> Set[str]: + env = jinja2.sandbox.ImmutableSandboxedEnvironment( + trim_blocks=True, + lstrip_blocks=True, + extensions=[AssistantTracker, jinja2.ext.loopcontrols], + ) + parsed_content = env.parse(chat_template) + template_vars = jinja2.meta.find_undeclared_variables(parsed_content) + return template_vars + + +_cached_resolve_chat_template_kwargs = lru_cache(_resolve_chat_template_kwargs) + + +@lru_cache +def _get_hf_base_chat_template_params() -> frozenset[str]: + from transformers import PreTrainedTokenizer + + # Get standard parameters from HuggingFace's base tokenizer class. + # This dynamically extracts parameters from PreTrainedTokenizer's + # apply_chat_template method, ensuring compatibility with tokenizers + # that use **kwargs to receive standard parameters. + + # Read signature from HF's base class - the single source of truth + base_sig = inspect.signature(PreTrainedTokenizer.apply_chat_template) + + # Exclude VAR_KEYWORD (**kwargs) and VAR_POSITIONAL (*args) placeholders + return frozenset( + p.name + for p in base_sig.parameters.values() + if p.kind + not in (inspect.Parameter.VAR_KEYWORD, inspect.Parameter.VAR_POSITIONAL) + ) + + +def resolve_chat_template_kwargs( + tokenizer: HfTokenizer, + chat_template: str, + chat_template_kwargs: dict[str, Any], + raise_on_unexpected: bool = True, +) -> dict[str, Any]: + # We exclude chat_template from kwargs here, because + # chat template has been already resolved at this stage + unexpected_vars = {"chat_template", "tokenize"} + if raise_on_unexpected and ( + unexpected_in_kwargs := unexpected_vars & chat_template_kwargs.keys() + ): + raise ValueError( + "Found unexpected chat template kwargs from request: " + f"{unexpected_in_kwargs}" + ) + + fn_kw = { + k + for k in chat_template_kwargs + if supports_kw(tokenizer.apply_chat_template, k, allow_var_kwargs=False) + } + template_vars = _cached_resolve_chat_template_kwargs(chat_template) + + # Allow standard HF parameters even if tokenizer uses **kwargs to receive them + hf_base_params = _get_hf_base_chat_template_params() + + accept_vars = (fn_kw | template_vars | hf_base_params) - unexpected_vars + return {k: v for k, v in chat_template_kwargs.items() if k in accept_vars} + + +def safe_apply_chat_template( + model_config: "ModelConfig", + tokenizer: HfTokenizer, + conversation: list[ConversationMessage], + *, + tools: list[dict[str, Any]] | None = None, + chat_template: str | None = None, + tokenize: bool = True, + **kwargs, +) -> str | list[int]: + chat_template = resolve_chat_template( + tokenizer, + chat_template=chat_template, + tools=tools, + model_config=model_config, + ) + if chat_template is None: + raise ChatTemplateResolutionError( + "As of transformers v4.44, default chat template is no longer " + "allowed, so you must provide a chat template if the tokenizer " + "does not define one." + ) + + resolved_kwargs = resolve_chat_template_kwargs( + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=kwargs, + ) + + try: + return tokenizer.apply_chat_template( + conversation=conversation, # type: ignore[arg-type] + tools=tools, # type: ignore[arg-type] + chat_template=chat_template, + tokenize=tokenize, + **resolved_kwargs, + ) + # External library exceptions can sometimes occur despite the framework's + # internal exception management capabilities. + except Exception as e: + # Log and report any library-related exceptions for further + # investigation. + logger.exception( + "An error occurred in `transformers` while applying chat template" + ) + raise ValueError(str(e)) from e + + +class HfRenderer(RendererLike): + @classmethod + def from_config( + cls, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + return cls(config, tokenizer_kwargs) + + def __init__( + self, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> None: + super().__init__() + + self.config = config + + if config.skip_tokenizer_init: + tokenizer = None + else: + tokenizer = cast( + HfTokenizer, + cached_get_tokenizer( + tokenizer_cls=CachedHfTokenizer, # type: ignore[type-abstract] + **tokenizer_kwargs, + ), + ) + + self._tokenizer = tokenizer + + @property + def tokenizer(self) -> HfTokenizer | None: + return self._tokenizer + + def get_tokenizer(self) -> HfTokenizer: + tokenizer = self.tokenizer + if tokenizer is None: + raise ValueError("Tokenizer not available when `skip_tokenizer_init=True`") + + return tokenizer + + def render_messages( + self, + messages: list[ChatCompletionMessageParam], + chat_template_content_format: ChatTemplateContentFormatOption = "auto", + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + model_config = self.config + tokenizer = self.get_tokenizer() + + conversation, mm_data, mm_uuids = parse_chat_messages( + messages, + model_config, + content_format=resolve_chat_template_content_format( + chat_template=kwargs.get("chat_template"), + tools=kwargs.get("tools"), + given_format=chat_template_content_format, + tokenizer=tokenizer, + model_config=model_config, + ), + ) + + prompt_raw = safe_apply_chat_template( + model_config, + tokenizer, + conversation, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] + + async def render_messages_async( + self, + messages: list[ChatCompletionMessageParam], + chat_template_content_format: ChatTemplateContentFormatOption = "auto", + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + model_config = self.config + tokenizer = self.get_tokenizer() + + conversation, mm_data, mm_uuids = await parse_chat_messages_async( + messages, + model_config, + content_format=resolve_chat_template_content_format( + chat_template=kwargs.get("chat_template"), + tools=kwargs.get("tools"), + given_format=chat_template_content_format, + tokenizer=tokenizer, + model_config=model_config, + ), + ) + + prompt_raw = safe_apply_chat_template( + model_config, + tokenizer, + conversation, + **kwargs, + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] diff --git a/vllm/renderers/mistral.py b/vllm/renderers/mistral.py new file mode 100644 index 00000000000..c45fb1f77ed --- /dev/null +++ b/vllm/renderers/mistral.py @@ -0,0 +1,147 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ConversationMessage, + parse_chat_messages, + parse_chat_messages_async, +) +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.logger import init_logger +from vllm.tokenizers import cached_get_tokenizer +from vllm.tokenizers.mistral import MistralTokenizer +from vllm.utils.async_utils import make_async + +from .protocol import RendererLike + +logger = init_logger(__name__) + + +def safe_apply_chat_template( + tokenizer: MistralTokenizer, + messages: list[ChatCompletionMessageParam], + **kwargs, +) -> str | list[int]: + from mistral_common.exceptions import MistralCommonException + + try: + return tokenizer.apply_chat_template(messages, **kwargs) + # mistral-common uses assert statements to stop processing of input + # if input does not comply with the expected format. + # We convert those assertion errors to ValueErrors so they can be + # properly caught in the preprocessing_input step + except (AssertionError, MistralCommonException) as e: + raise ValueError(str(e)) from e + + # External library exceptions can sometimes occur despite the framework's + # internal exception management capabilities. + except Exception as e: + # Log and report any library-related exceptions for further + # investigation. + logger.exception( + "An error occurred in `mistral_common` while applying chat template" + ) + raise ValueError(str(e)) from e + + +class MistralRenderer(RendererLike): + @classmethod + def from_config( + cls, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + return cls(config, tokenizer_kwargs) + + def __init__( + self, + config: ModelConfig, + tokenizer_kwargs: dict[str, Any], + ) -> None: + super().__init__() + + self.config = config + + if config.skip_tokenizer_init: + tokenizer = None + else: + tokenizer = cached_get_tokenizer( + tokenizer_cls=MistralTokenizer, + **tokenizer_kwargs, + ) + + self._tokenizer = tokenizer + + self._apply_chat_template_executor = ThreadPoolExecutor(max_workers=1) + self._apply_chat_template_async = make_async( + safe_apply_chat_template, executor=self._apply_chat_template_executor + ) + + @property + def tokenizer(self) -> MistralTokenizer | None: + return self._tokenizer + + def get_tokenizer(self) -> MistralTokenizer: + tokenizer = self.tokenizer + if tokenizer is None: + raise ValueError("Tokenizer not available when `skip_tokenizer_init=True`") + + return tokenizer + + def render_messages( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = parse_chat_messages( + messages, + self.config, + content_format="string", + ) + + prompt_raw = safe_apply_chat_template(tokenizer, messages, **kwargs) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] + + async def render_messages_async( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + tokenizer = self.get_tokenizer() + conversation, mm_data, mm_uuids = await parse_chat_messages_async( + messages, + self.config, + content_format="string", + ) + + prompt_raw = await self._apply_chat_template_async( + tokenizer, messages, **kwargs + ) + + prompt = ( + TextPrompt(prompt=prompt_raw) + if isinstance(prompt_raw, str) + else TokensPrompt(prompt_token_ids=prompt_raw) + ) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt # type: ignore[return-value] diff --git a/vllm/renderers/protocol.py b/vllm/renderers/protocol.py new file mode 100644 index 00000000000..e788f431b0f --- /dev/null +++ b/vllm/renderers/protocol.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import TYPE_CHECKING, Any, Protocol + +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.tokenizers import TokenizerLike + +if TYPE_CHECKING: + from vllm.config import ModelConfig + from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ConversationMessage, + ) + + +class RendererLike(Protocol): + @classmethod + def from_config( + cls, + config: "ModelConfig", + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + raise NotImplementedError + + @property + def tokenizer(self) -> TokenizerLike | None: + raise NotImplementedError + + def get_tokenizer(self) -> TokenizerLike: + tokenizer = self.tokenizer + if tokenizer is None: + raise ValueError("Tokenizer not available when `skip_tokenizer_init=True`") + + return tokenizer + + def render_messages( + self, + messages: list["ChatCompletionMessageParam"], + **kwargs, + ) -> tuple[list["ConversationMessage"], TextPrompt | TokensPrompt]: + raise NotImplementedError + + async def render_messages_async( + self, + messages: list["ChatCompletionMessageParam"], + **kwargs, + ) -> tuple[list["ConversationMessage"], TextPrompt | TokensPrompt]: + return self.render_messages(messages, **kwargs) diff --git a/vllm/renderers/registry.py b/vllm/renderers/registry.py new file mode 100644 index 00000000000..5269978b5b2 --- /dev/null +++ b/vllm/renderers/registry.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any + +from vllm.logger import init_logger +from vllm.tokenizers.registry import tokenizer_args_from_config +from vllm.utils.import_utils import resolve_obj_by_qualname + +from .protocol import RendererLike + +if TYPE_CHECKING: + from vllm.config import ModelConfig + +logger = init_logger(__name__) + + +_VLLM_RENDERERS = { + "deepseek_v32": ("deepseek_v32", "DeepseekV32Renderer"), + "hf": ("hf", "HfRenderer"), + "grok2": ("grok2", "Grok2Renderer"), + "mistral": ("mistral", "MistralRenderer"), + "terratorch": ("terratorch", "TerratorchRenderer"), +} + + +@dataclass +class RendererRegistry: + # Renderer mode -> (renderer module, renderer class) + renderers: dict[str, tuple[str, str]] = field(default_factory=dict) + + def register(self, renderer_mode: str, module: str, class_name: str) -> None: + if renderer_mode in self.renderers: + logger.warning( + "%s.%s is already registered for renderer_mode=%r. " + "It is overwritten by the new one.", + module, + class_name, + renderer_mode, + ) + + self.renderers[renderer_mode] = (module, class_name) + + return None + + def load_renderer_cls(self, renderer_mode: str) -> type[RendererLike]: + if renderer_mode not in self.renderers: + raise ValueError(f"No renderer registered for {renderer_mode=!r}.") + + module, class_name = self.renderers[renderer_mode] + logger.debug_once(f"Loading {class_name} for {renderer_mode=!r}") + + return resolve_obj_by_qualname(f"{module}.{class_name}") + + def load_renderer( + self, + renderer_mode: str, + config: "ModelConfig", + tokenizer_kwargs: dict[str, Any], + ) -> RendererLike: + renderer_cls = self.load_renderer_cls(renderer_mode) + return renderer_cls.from_config(config, tokenizer_kwargs) + + +RENDERER_REGISTRY = RendererRegistry( + { + mode: (f"vllm.renderers.{mod_relname}", cls_name) + for mode, (mod_relname, cls_name) in _VLLM_RENDERERS.items() + } +) +"""The global `RendererRegistry` instance.""" + + +def renderer_from_config(config: "ModelConfig", **kwargs): + tokenizer_mode, tokenizer_name, args, kwargs = tokenizer_args_from_config( + config, **kwargs + ) + + if config.tokenizer_mode == "auto" and config.model_impl == "terratorch": + renderer_mode = "terratorch" + else: + renderer_mode = tokenizer_mode + + return RENDERER_REGISTRY.load_renderer( + renderer_mode, + config, + tokenizer_kwargs={**kwargs, "tokenizer_name": tokenizer_name}, + ) diff --git a/vllm/renderers/terratorch.py b/vllm/renderers/terratorch.py new file mode 100644 index 00000000000..fc41a94c85b --- /dev/null +++ b/vllm/renderers/terratorch.py @@ -0,0 +1,85 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +from vllm.config import ModelConfig +from vllm.entrypoints.chat_utils import ( + ChatCompletionMessageParam, + ConversationMessage, + parse_chat_messages, + parse_chat_messages_async, +) +from vllm.inputs import TextPrompt, TokensPrompt +from vllm.logger import init_logger +from vllm.tokenizers import TokenizerLike + +from .protocol import RendererLike + +logger = init_logger(__name__) + + +class TerratorchRenderer(RendererLike): + @classmethod + def from_config( + cls, + config: "ModelConfig", + tokenizer_kwargs: dict[str, Any], + ) -> "RendererLike": + return cls(config) + + def __init__(self, config: ModelConfig) -> None: + super().__init__() + + self.config = config + + if not config.skip_tokenizer_init: + raise ValueError("Terratorch renderer requires `skip_tokenizer_init=True`") + + @property + def tokenizer(self) -> TokenizerLike | None: + return None + + def get_tokenizer(self) -> TokenizerLike: + raise ValueError("Tokenizer not available for Terratorch renderer") + + def render_messages( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + model_config = self.config + + conversation, mm_data, mm_uuids = parse_chat_messages( + messages, + model_config, + content_format="string", + ) + + prompt = TokensPrompt(prompt_token_ids=[1]) + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt + + async def render_messages_async( + self, + messages: list[ChatCompletionMessageParam], + **kwargs, + ) -> tuple[list[ConversationMessage], TextPrompt | TokensPrompt]: + model_config = self.config + + conversation, mm_data, mm_uuids = await parse_chat_messages_async( + messages, + model_config, + content_format="string", + ) + + prompt = TokensPrompt(prompt_token_ids=[1]) # Dummy token IDs + if mm_data is not None: + prompt["multi_modal_data"] = mm_data + if mm_uuids is not None: + prompt["multi_modal_uuids"] = mm_uuids + + return conversation, prompt diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index cceb51796a4..4f1126d1720 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -23,9 +23,10 @@ from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry from vllm.outputs import PoolingRequestOutput, RequestOutput from vllm.plugins.io_processors import get_io_processor from vllm.pooling_params import PoolingParams +from vllm.renderers import RendererLike from vllm.sampling_params import SamplingParams from vllm.tasks import SupportedTask -from vllm.tokenizers import TokenizerLike, cached_tokenizer_from_config +from vllm.tokenizers import TokenizerLike from vllm.tracing import init_tracer from vllm.transformers_utils.config import maybe_register_config_serialize_by_value from vllm.usage.usage_lib import UsageContext @@ -106,9 +107,7 @@ class AsyncLLM(EngineClient): "enabling logging without default stat loggers." ) - tokenizer = cached_tokenizer_from_config(self.model_config) - - self.input_processor = InputProcessor(self.vllm_config, tokenizer) + self.input_processor = InputProcessor(self.vllm_config) self.io_processor = get_io_processor( self.vllm_config, self.model_config.io_processor_plugin, @@ -709,13 +708,12 @@ class AsyncLLM(EngineClient): def tokenizer(self) -> TokenizerLike | None: return self.input_processor.tokenizer - async def get_tokenizer(self) -> TokenizerLike: - if self.tokenizer is None: - raise ValueError( - "Unable to get tokenizer because `skip_tokenizer_init=True`" - ) + def get_tokenizer(self) -> TokenizerLike: + return self.input_processor.get_tokenizer() - return self.tokenizer + @property + def renderer(self) -> RendererLike: + return self.input_processor.renderer async def is_tracing_enabled(self) -> bool: return self.observability_config.otlp_traces_endpoint is not None # type: ignore diff --git a/vllm/v1/engine/input_processor.py b/vllm/v1/engine/input_processor.py index 06a8c4b69f1..4d5f1dca6a1 100644 --- a/vllm/v1/engine/input_processor.py +++ b/vllm/v1/engine/input_processor.py @@ -19,6 +19,7 @@ from vllm.multimodal.parse import MultiModalDataParser from vllm.multimodal.processing.context import set_request_id from vllm.multimodal.utils import argsort_mm_positions from vllm.pooling_params import PoolingParams +from vllm.renderers import RendererLike from vllm.sampling_params import _SAMPLING_EPS, SamplingParams from vllm.tokenizers import TokenizerLike from vllm.tokenizers.mistral import MistralTokenizer @@ -45,7 +46,6 @@ class InputProcessor: def __init__( self, vllm_config: VllmConfig, - tokenizer: TokenizerLike | None, mm_registry: MultiModalRegistry = MULTIMODAL_REGISTRY, ) -> None: self.vllm_config = vllm_config @@ -61,8 +61,7 @@ class InputProcessor: self.input_preprocessor = InputPreprocessor( self.model_config, - tokenizer, - self.vllm_config.observability_config, + vllm_config.observability_config, mm_registry, mm_processor_cache=self.mm_processor_cache, ) @@ -71,6 +70,13 @@ class InputProcessor: def tokenizer(self) -> TokenizerLike | None: return self.input_preprocessor.tokenizer + def get_tokenizer(self) -> TokenizerLike: + return self.input_preprocessor.get_tokenizer() + + @property + def renderer(self) -> RendererLike: + return self.input_preprocessor.renderer + def _validate_logprobs( self, params: SamplingParams, diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py index 78eeb70f116..5811e94dd3c 100644 --- a/vllm/v1/engine/llm_engine.py +++ b/vllm/v1/engine/llm_engine.py @@ -21,9 +21,10 @@ from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry from vllm.outputs import PoolingRequestOutput, RequestOutput from vllm.plugins.io_processors import get_io_processor from vllm.pooling_params import PoolingParams +from vllm.renderers import RendererLike from vllm.sampling_params import SamplingParams from vllm.tasks import SupportedTask -from vllm.tokenizers import TokenizerLike, cached_tokenizer_from_config +from vllm.tokenizers import TokenizerLike from vllm.tracing import init_tracer from vllm.usage.usage_lib import UsageContext from vllm.v1.engine import EngineCoreRequest @@ -84,9 +85,7 @@ class LLMEngine: self.dp_group = None self.should_execute_dummy_batch = False - tokenizer = cached_tokenizer_from_config(self.model_config) - - self.input_processor = InputProcessor(self.vllm_config, tokenizer) + self.input_processor = InputProcessor(self.vllm_config) self.io_processor = get_io_processor( self.vllm_config, self.model_config.io_processor_plugin, @@ -357,12 +356,11 @@ class LLMEngine: return self.input_processor.tokenizer def get_tokenizer(self) -> TokenizerLike: - if self.tokenizer is None: - raise ValueError( - "Unable to get tokenizer because `skip_tokenizer_init=True`" - ) + return self.input_processor.get_tokenizer() - return self.tokenizer + @property + def renderer(self) -> RendererLike: + return self.input_processor.renderer def do_log_stats(self) -> None: """Log stats if logging is enabled."""