Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions tensorrt_llm/llmapi/disagg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
143 changes: 143 additions & 0 deletions tensorrt_llm/serve/chat_tokenization.py
Original file line number Diff line number Diff line change
@@ -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
36 changes: 10 additions & 26 deletions tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Comment thread
coderabbitai[bot] marked this conversation as resolved.

# Get harmony stop tokens
harmony_stop_tokens = self.harmony_adapter.get_stop_tokens()
Expand Down
14 changes: 10 additions & 4 deletions tensorrt_llm/serve/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
74 changes: 44 additions & 30 deletions tensorrt_llm/serve/router_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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(
Expand All @@ -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)
Expand Down
11 changes: 11 additions & 0 deletions tests/unittest/disaggregated/test_disagg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading
Loading