diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index a8272b494645..1796051a29c9 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -311,6 +311,11 @@ def extract_router_config(server_cfg: dict) -> RouterConfig: args = server_cfg.pop("router", {}) router_type = args.pop("type", "round_robin") + if router_type == "kv_cache_aware" and "model_path" not in args: + model_path = server_cfg.get("model") + if model_path is not None: + args["model_path"] = model_path + # add fields that are not specific to router extract_keys = ["max_batch_size", "max_num_tokens"] for key in extract_keys: diff --git a/tensorrt_llm/serve/chat_tokenization.py b/tensorrt_llm/serve/chat_tokenization.py new file mode 100644 index 000000000000..d839e79bfff6 --- /dev/null +++ b/tensorrt_llm/serve/chat_tokenization.py @@ -0,0 +1,143 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import os +from typing import TYPE_CHECKING, Callable, Optional, cast + +from transformers import PretrainedConfig + +from tensorrt_llm.serve.openai_protocol import ChatCompletionRequest + +if TYPE_CHECKING: + from tensorrt_llm.serve.harmony_adapter import HarmonyAdapter + +ToolDict = dict[str, object] + + +def resolve_model_type_from_config(model_name_or_path: str) -> Optional[str]: + """Return the checkpoint's declared model type from its config metadata.""" + config_dict, _ = PretrainedConfig.get_config_dict(model_name_or_path) + model_type = config_dict.get("model_type") + return model_type if isinstance(model_type, str) else None + + +def uses_harmony_tokenization( + use_harmony: Optional[bool] = None, + model_type: Optional[str] = None, + model_type_resolver: Optional[Callable[[], Optional[str]]] = None, +) -> bool: + if os.getenv("DISABLE_HARMONY_ADAPTER", "0") == "1": + return False + if use_harmony is not None: + return use_harmony + if model_type is None and model_type_resolver is not None: + model_type = model_type_resolver() + return model_type == "gpt_oss" + + +def get_chat_completion_tool_dicts( + request: ChatCompletionRequest, empty_as_none: bool = False +) -> Optional[list[ToolDict]]: + if request.tools is None or (empty_as_none and not request.tools): + return None + tools: list[ToolDict] = [] + for tool in request.tools: + if hasattr(tool, "model_dump"): + tools.append(cast(ToolDict, tool.model_dump())) + elif isinstance(tool, dict): + tools.append(cast(ToolDict, tool)) + else: + raise TypeError(f"Unsupported tool type: {type(tool).__name__}") + return tools + + +def tokenize_harmony_chat_request( + request: ChatCompletionRequest, + harmony_adapter: Optional["HarmonyAdapter"] = None, + set_prompt_token_ids: bool = False, +) -> list[int]: + if request.prompt_token_ids is not None: + return request.prompt_token_ids + + from tensorrt_llm.serve import harmony_adapter as harmony_adapter_module + + adapter = harmony_adapter or harmony_adapter_module.get_harmony_adapter() + result = adapter.openai_to_harmony_tokens( + request.messages, + get_chat_completion_tool_dicts(request, empty_as_none=True), + reasoning_effort=harmony_adapter_module.maybe_transform_reasoning_effort( + request.reasoning_effort + ), + tool_choice=request.tool_choice, + ) + if set_prompt_token_ids: + request.prompt_token_ids = result + return result + + +def render_chat_request_for_tokenizer( + request: ChatCompletionRequest, tokenizer: object +) -> str | list[int]: + chat_template_kwargs = ( + dict(request.chat_template_kwargs) if getattr(request, "chat_template_kwargs", None) else {} + ) + chat_template_kwargs["tools"] = get_chat_completion_tool_dicts(request) + chat_template_kwargs["documents"] = request.documents + if request.chat_template is not None: + chat_template_kwargs["chat_template"] = request.chat_template + rendered = tokenizer.apply_chat_template( + [msg if isinstance(msg, dict) else dict(msg) for msg in request.messages], + add_generation_prompt=request.add_generation_prompt, + tokenize=False, + return_dict=False, + **chat_template_kwargs, + ) + if isinstance(rendered, str): + return rendered + return list(rendered) + + +def tokenize_chat_request_for_serving( + request: ChatCompletionRequest, + tokenizer_factory: Callable[[], object], + encode_rendered: Callable[[str, object], list[int]], + use_harmony: Optional[bool] = None, + model_type: Optional[str] = None, + model_type_resolver: Optional[Callable[[], Optional[str]]] = None, + harmony_adapter: Optional["HarmonyAdapter"] = None, + set_prompt_token_ids: bool = True, +) -> list[int]: + if request.prompt_token_ids is not None: + return request.prompt_token_ids + + if uses_harmony_tokenization( + use_harmony=use_harmony, + model_type=model_type, + model_type_resolver=model_type_resolver, + ): + return tokenize_harmony_chat_request( + request, + harmony_adapter=harmony_adapter, + set_prompt_token_ids=set_prompt_token_ids, + ) + + tokenizer = tokenizer_factory() + rendered = render_chat_request_for_tokenizer(request, tokenizer) + result = encode_rendered(rendered, tokenizer) if isinstance(rendered, str) else rendered + if set_prompt_token_ids: + request.prompt_token_ids = result + return result diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 153b35b5f6cb..c60bbd0449d3 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -56,6 +56,7 @@ from tensorrt_llm.runtime.kv_cache_hash import \ get_effective_kv_cache_event_hash_algo from tensorrt_llm.sampling_params import GuidedDecodingParams, SamplingParams +from tensorrt_llm.serve.chat_tokenization import tokenize_harmony_chat_request from tensorrt_llm.serve.chat_utils import (load_chat_template, parse_chat_messages_coroutines, resolve_top_level_model_type) @@ -100,8 +101,7 @@ from tensorrt_llm.visual_gen import VisualGen from .._utils import nvtx_mark, set_prometheus_multiproc_dir -from .harmony_adapter import (HarmonyAdapter, get_harmony_adapter, - maybe_transform_reasoning_effort) +from .harmony_adapter import HarmonyAdapter, get_harmony_adapter # yapf: enable @@ -2037,34 +2037,18 @@ async def create_streaming_generator(promise: RequestOutput, # NOTE: WAR for Disagg failure, may affect perf if no warmup if not self.harmony_adapter: self.harmony_adapter = get_harmony_adapter() - # Convert Pydantic models to dictionaries for JSON serialization (standard pattern) - tools_dict = None - if request.tools: - tools_dict = [tool.model_dump() for tool in request.tools] - - # Reasoning effort precedence: request.reasoning_effort > system message parsing > serving default - reasoning_effort = maybe_transform_reasoning_effort( - request.reasoning_effort) - # Get tool_choice from request - tool_choice = getattr(request, 'tool_choice', None) # Reuse pre-tokenized harmony tokens when forwarded by an upstream # context worker (disaggregated serving). Otherwise, run the # Harmony adapter on the request messages. - if request.prompt_token_ids is not None: - harmony_tokens = request.prompt_token_ids - else: - try: - harmony_tokens = self.harmony_adapter.openai_to_harmony_tokens( - request.messages, - tools_dict, - reasoning_effort=reasoning_effort, - tool_choice=tool_choice) - except Exception: - logger.error(f"messages_dict: {request.messages}") - logger.error(f"tools_dict: {tools_dict}") - logger.error(f"request: {request}") - raise + try: + harmony_tokens = tokenize_harmony_chat_request( + request, harmony_adapter=self.harmony_adapter) + except Exception: + logger.error("messages_dict: %s", request.messages) + logger.error("tools: %s", request.tools) + logger.error("request: %s", request) + raise # Get harmony stop tokens harmony_stop_tokens = self.harmony_adapter.get_stop_tokens() diff --git a/tensorrt_llm/serve/router.py b/tensorrt_llm/serve/router.py index f0588a7e82f6..bb7ad601d732 100644 --- a/tensorrt_llm/serve/router.py +++ b/tensorrt_llm/serve/router.py @@ -798,14 +798,19 @@ def __init__(self, tokens_per_block: Optional[int] = None, custom_tokenizer: Optional[str] = None, tokenizer_dir: Optional[str] = None, + use_harmony: Optional[bool] = None, + model_path: Optional[str] = None, track_routed_blocks: bool = True, load_weight: float = 0.25, load_cap: float = float("inf"), - **kwargs): + **kwargs) -> None: super().__init__(server_role, servers, metadata_server_cfg, metadata_server, **kwargs) - self._init_block_hashing(tokens_per_block, custom_tokenizer, - tokenizer_dir) + self._init_block_hashing(tokens_per_block, + custom_tokenizer, + tokenizer_dir, + use_harmony=use_harmony, + model_path=model_path) self._init_load_balancing(servers, use_tokens) # TODO: use max_num_tokens? per server? self._max_batch_size = max_batch_size @@ -1252,11 +1257,12 @@ def __init__(self, use_token_ids: bool = False, hash_skip_count: int = 0, max_sessions: int = 100000, + use_harmony: Optional[bool] = None, **kwargs): super().__init__(server_role, servers, metadata_server_cfg, metadata_server, **kwargs) self._init_load_balancing(servers) - self._init_block_hashing(tokens_per_block) + self._init_block_hashing(tokens_per_block, use_harmony=use_harmony) self._match_threshold = match_threshold self._use_token_ids = use_token_ids diff --git a/tensorrt_llm/serve/router_utils.py b/tensorrt_llm/serve/router_utils.py index a96d087acf06..5297f685fd29 100644 --- a/tensorrt_llm/serve/router_utils.py +++ b/tensorrt_llm/serve/router_utils.py @@ -28,6 +28,10 @@ from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import Block as V2Block from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import ReuseScope from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import RootBlock as V2RootBlock +from tensorrt_llm.serve.chat_tokenization import ( + resolve_model_type_from_config, + tokenize_chat_request_for_serving, +) from tensorrt_llm.serve.openai_protocol import ChatCompletionRequest, CompletionRequest KV_CACHE_HASH_ALGO_DEFAULT = kv_cache_hash.KV_CACHE_HASH_ALGO_DEFAULT @@ -188,19 +192,26 @@ def _init_block_hashing( tokens_per_block: Optional[int] = None, custom_tokenizer: Optional[str] = None, tokenizer_dir: Optional[str] = None, - ): + use_harmony: Optional[bool] = None, + model_path: Optional[str] = None, + ) -> None: env_tokens_per_block = os.environ.get("TRTLLM_KVCACHE_AWARE_ROUTER_HASH_TOKENS_PER_BLOCK") if env_tokens_per_block is not None: tokens_per_block = int(env_tokens_per_block) self._tpb_auto = tokens_per_block is None self._tokens_per_block = 32 if tokens_per_block is None else tokens_per_block self._tokenizers: dict = {} + self._model_types: dict[str, Optional[str]] = {} self._custom_tokenizer = custom_tokenizer self._tokenizer_dir = tokenizer_dir + self._model_path = model_path + self._use_harmony = use_harmony logger.info( f"BlockHashMixin: tokens_per_block={self._tokens_per_block}" f"{' (auto, adopts worker)' if self._tpb_auto else ''}" f", custom_tokenizer={self._custom_tokenizer}" + f", model_path={self._model_path}" + f", use_harmony={self._use_harmony}" ) def _get_tokenizer(self, model: str): @@ -236,34 +247,31 @@ def _encode_with_prefix_cache(self, rendered: str, key: int, tokenizer) -> list[ cache.popitem(last=False) return ids + def _get_model_type(self) -> Optional[str]: + model_path = self._model_path or self._tokenizer_dir + if model_path is None: + return None + if model_path not in self._model_types: + try: + self._model_types[model_path] = resolve_model_type_from_config(model_path) + except (OSError, ValueError) as error: + logger.warning( + "Unable to resolve model type from checkpoint config at %s: %s. " + "Set use_harmony explicitly if the checkpoint uses Harmony.", + model_path, + error, + ) + self._model_types[model_path] = None + return self._model_types[model_path] + def _tokenize(self, request: OpenAIRequest) -> list[list[int]]: # Handle ChatCompletionRequest (has messages, not prompt) if isinstance(request, ChatCompletionRequest): - if request.prompt_token_ids is not None: - return [request.prompt_token_ids] - tokenizer = self._get_tokenizer(request.model) - tool_dicts = ( - None - if getattr(request, "tools", None) is None - else [ - tool.model_dump() if hasattr(tool, "model_dump") else tool - for tool in request.tools - ] - ) - chat_template_kwargs = ( - request.chat_template_kwargs - if getattr(request, "chat_template_kwargs", None) - else {} - ) - rendered = tokenizer.apply_chat_template( - [msg if isinstance(msg, dict) else dict(msg) for msg in request.messages], - add_generation_prompt=request.add_generation_prompt, - tokenize=False, - return_dict=False, - tools=tool_dicts, - **chat_template_kwargs, - ) - if isinstance(rendered, str): + + def tokenizer_factory() -> object: + return self._get_tokenizer(request.model) + + def encode_rendered(rendered: str, tokenizer: object) -> list[int]: key = hash( "".join( str( @@ -274,10 +282,16 @@ def _tokenize(self, request: OpenAIRequest) -> list[list[int]]: for msg in request.messages[:2] ) ) - result = self._encode_with_prefix_cache(rendered, key, tokenizer) - else: - result = list(rendered) - request.prompt_token_ids = result + return self._encode_with_prefix_cache(rendered, key, tokenizer) + + result = tokenize_chat_request_for_serving( + request, + tokenizer_factory=tokenizer_factory, + encode_rendered=encode_rendered, + use_harmony=self._use_harmony, + model_type_resolver=self._get_model_type, + set_prompt_token_ids=True, + ) return [result] # Handle CompletionRequest (has prompt) diff --git a/tests/unittest/disaggregated/test_disagg_utils.py b/tests/unittest/disaggregated/test_disagg_utils.py index f3556bae1fe5..1cc169ef662c 100644 --- a/tests/unittest/disaggregated/test_disagg_utils.py +++ b/tests/unittest/disaggregated/test_disagg_utils.py @@ -195,6 +195,17 @@ def test_extract_router_config_propagates_tokens_per_block(): }).args +def test_extract_router_config_propagates_kv_model_path() -> None: + cfg = { + "model": "/models/gpt-oss-checkpoint", + "router": { + "type": "kv_cache_aware" + }, + } + router_config = extract_router_config(cfg) + assert router_config.args["model_path"] == "/models/gpt-oss-checkpoint" + + def test_get_server_configs_dict(): server_configs = [ CtxGenServerConfig(type="ctx", diff --git a/tests/unittest/disaggregated/test_router.py b/tests/unittest/disaggregated/test_router.py index ddaf12a0c8e2..1eaaaa9ea82a 100644 --- a/tests/unittest/disaggregated/test_router.py +++ b/tests/unittest/disaggregated/test_router.py @@ -2,6 +2,8 @@ import copy import random import threading +from pathlib import Path +from types import SimpleNamespace from unittest import mock import aiohttp @@ -1861,6 +1863,169 @@ def _mock_tokenizer(token_ids=None): return tok +def test_router_model_type_uses_checkpoint_config(tmp_path: Path) -> None: + (tmp_path / "config.json").write_text('{"model_type": "gpt_oss"}', + encoding="utf-8") + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + model_path=str(tmp_path)) + + assert router._get_model_type() == "gpt_oss" + + +@pytest.mark.asyncio +async def test_gpt_oss_router_tokens_match_chat_harmony_server_input() -> None: + """KV-cache routing must hash the same Harmony tokens used by the server.""" + from tensorrt_llm.serve.openai_server import OpenAIServer + + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32, + model_path="/models/gpt-oss-checkpoint") + router_tokenizer = _mock_tokenizer(token_ids=[900, 901, 902]) + harmony_tokens = [100, 101, 102, 103] + harmony_adapter = mock.MagicMock() + harmony_adapter.openai_to_harmony_tokens.return_value = harmony_tokens + harmony_adapter.get_stop_tokens.return_value = [42] + promise = mock.MagicMock() + promise.prompt_token_ids = [] + + request = ChatCompletionRequest( + model="my-model", + messages=[{ + "role": "developer", + "content": "Use tools when useful." + }, { + "role": "user", + "content": "weather in Paris?" + }], + tools=[_get_weather_tool()], + tool_choice="auto", + reasoning_effort="medium", + stream=True, + max_completion_tokens=1, + ) + router_request = copy.deepcopy(request) + server_request = copy.deepcopy(request) + + server = OpenAIServer.__new__(OpenAIServer) + server.allow_request_chat_template = False + server.await_disconnected = mock.AsyncMock() + server.generator = SimpleNamespace( + args=SimpleNamespace(num_postprocess_workers=0), + generate_async=mock.MagicMock(return_value=promise), + ) + server.harmony_adapter = harmony_adapter + server.model_config = SimpleNamespace(vocab_size=1000) + server.tokenizer = SimpleNamespace(tokenizer=SimpleNamespace( + vocab_size=1000)) + + with mock.patch.object( + router, "_get_tokenizer", + return_value=router_tokenizer), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "get_harmony_adapter", + return_value=harmony_adapter), mock.patch( + "tensorrt_llm.serve.router_utils." + "resolve_model_type_from_config", + return_value="gpt_oss") as resolve_model_type: + router_token_ids = router._tokenize(router_request)[0] + await server.chat_harmony(server_request, raw_request=None) + + server_token_ids = server.generator.generate_async.call_args.kwargs[ + "inputs"] + assert router_token_ids == server_token_ids + assert router_request.prompt_token_ids == harmony_tokens + first_call, second_call = harmony_adapter.openai_to_harmony_tokens.call_args_list + assert first_call.args == second_call.args + assert first_call.kwargs == second_call.kwargs + resolve_model_type.assert_called_once_with("/models/gpt-oss-checkpoint") + router_tokenizer.apply_chat_template.assert_not_called() + + +def test_gpt_oss_router_respects_disable_harmony_adapter( + monkeypatch: pytest.MonkeyPatch) -> None: + """Router follows the same DISABLE_HARMONY_ADAPTER gate as the server.""" + monkeypatch.setenv("DISABLE_HARMONY_ADAPTER", "1") + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32, + model_path="/models/gpt-oss-checkpoint") + router_tokenizer = _mock_tokenizer(token_ids=[900, 901, 902]) + harmony_adapter = mock.MagicMock() + + request = ChatCompletionRequest( + model="openai/gpt-oss-20b", + messages=[{ + "role": "user", + "content": "weather in Paris?" + }], + tools=[_get_weather_tool()], + ) + + with mock.patch.object( + router, "_get_tokenizer", + return_value=router_tokenizer), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "get_harmony_adapter", + return_value=harmony_adapter), mock.patch( + "tensorrt_llm.serve.router_utils." + "resolve_model_type_from_config", + side_effect=AssertionError( + "disabled Harmony must not load model config")): + assert router._tokenize(request) == [[900, 901, 902]] + + harmony_adapter.openai_to_harmony_tokens.assert_not_called() + router_tokenizer.apply_chat_template.assert_called_once() + + +@pytest.mark.asyncio +async def test_chat_harmony_preserves_original_tool_conversion_error() -> None: + """Harmony diagnostics must not rerun the conversion that already failed.""" + from tensorrt_llm.serve.openai_server import OpenAIServer + + original_error = RuntimeError("original tool conversion failure") + diagnostic_error = RuntimeError("diagnostic tool conversion failure") + + class FailingTool: + + def __init__(self) -> None: + self.calls = 0 + + def model_dump(self) -> dict[str, object]: + self.calls += 1 + if self.calls == 1: + raise original_error + raise diagnostic_error + + failing_tool = FailingTool() + request = ChatCompletionRequest( + model="my-model", + messages=[{ + "role": "user", + "content": "weather in Paris?" + }], + ) + object.__setattr__(request, "tools", [failing_tool]) + + server = OpenAIServer.__new__(OpenAIServer) + server.allow_request_chat_template = False + server.harmony_adapter = mock.MagicMock() + server.create_error_response = mock.MagicMock( + return_value=str(original_error)) + + response = await server.chat_harmony(request, raw_request=None) + + assert response == str(original_error) + assert failing_tool.calls == 1 + server.create_error_response.assert_called_once_with( + message=str(original_error), err_type="internal_error") + + @pytest.mark.parametrize("router_class", [KvCacheAwareRouter, ConversationRouter]) def test_tokenize_forwards_tools_and_chat_template_kwargs(router_class): @@ -1880,6 +2045,8 @@ def test_tokenize_forwards_tools_and_chat_template_kwargs(router_class): tokens_per_block=32) tok = _mock_tokenizer() + documents = [{"title": "Paris", "text": "Paris is in France."}] + chat_template = "{% for message in messages %}{{ message.content }}{% endfor %}" with mock.patch.object(router, "_get_tokenizer", return_value=tok): req = ChatCompletionRequest( model="TinyLlama", @@ -1888,6 +2055,8 @@ def test_tokenize_forwards_tools_and_chat_template_kwargs(router_class): "content": "what's the weather in Paris?" }], tools=[_get_weather_tool()], + documents=documents, + chat_template=chat_template, chat_template_kwargs={"thinking": True}, ) router._tokenize(req) @@ -1904,6 +2073,8 @@ def test_tokenize_forwards_tools_and_chat_template_kwargs(router_class): assert "parameters" in tool_dict["function"] # chat_template_kwargs must be forwarded as **kwargs (not nested). assert kwargs.get("thinking") is True + assert kwargs["documents"] == documents + assert kwargs["chat_template"] == chat_template @pytest.mark.parametrize("router_class",