Skip to content
Closed
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
149 changes: 145 additions & 4 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -703,6 +703,10 @@ def __init__(
# responses and flushing them at a synchronised point in the executor
# loop avoids the mismatch.
self._pending_transfer_responses: List[Tuple[int, LlmResponse]] = []
self._disagg_generation_trans_in_progress_requests: Dict[
int, LlmRequest] = {}
self._pending_disagg_generation_trans_complete_requests: Dict[
int, LlmRequest] = {}
# Same buffer-then-synced-flush pattern as _pending_transfer_responses
# above: _handle_responses and _append_iter_stats are reached from
# per-rank-divergent gates, so their tp_allgather collectives are
Expand Down Expand Up @@ -2413,6 +2417,8 @@ def _pp_schedule_and_propagate(self, microbatch_id: int):
serializable_schedule = None
wait_for_disagg_gen_transfer_progress = False
is_dp_broadcast = self.dist.tp_size > 1 and self.enable_attention_dp
if getattr(self, "kv_cache_transceiver", None):
self._sync_pending_disagg_generation_trans_complete_draft_tokens()
if self.dist.rank == 0 or (self.dist.is_first_pp_rank
and is_dp_broadcast):
scheduled_batch, fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule(
Expand Down Expand Up @@ -3355,6 +3361,120 @@ def _commit_kv_cache_stats(self,
self.kv_cache_manager.commit_scheduled_kv_cache_stats(
scheduled_batch)

@staticmethod

@Shixiaowei02 Shixiaowei02 Jul 21, 2026

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The draft-token sync itself is small, but it adds ~10 helpers, two tracking dicts, and 6 touch points in py_executor.py (our most safety-critical file) mainly to avoid an active_requests scan. Could check_gen_transfer_status return the completed request objects (as it already does for cancelled ones) so we can set the draft tokens inline on completion? That would let us drop the _disagg_generation_trans_in_progress_requests tracking layer (and the cancel-path leak it introduces) and keep the footprint much smaller. cc @chuangz0

def _get_disagg_generation_transfer_id(
request: LlmRequest) -> Optional[int]:
disagg_params = getattr(request, "py_disaggregated_params", None)
transfer_id = getattr(disagg_params, "disagg_request_id", None)
if isinstance(transfer_id, int):
return transfer_id
request_id = getattr(request, "py_request_id", None)
if isinstance(request_id, int):
return request_id
request_id = getattr(request, "request_id", None)
if isinstance(request_id, int):
return request_id
return None

def _get_disagg_generation_trans_in_progress_requests(
self) -> Dict[int, LlmRequest]:
requests = getattr(self,
"_disagg_generation_trans_in_progress_requests",
None)
if requests is None:
requests = {}
self._disagg_generation_trans_in_progress_requests = requests
return requests

def _get_pending_disagg_generation_trans_complete_requests(
self) -> Dict[int, LlmRequest]:
requests = getattr(
self, "_pending_disagg_generation_trans_complete_requests", None)
if requests is None:
requests = {}
self._pending_disagg_generation_trans_complete_requests = requests
return requests

def _track_disagg_generation_trans_in_progress(
self, requests: Iterable[LlmRequest]) -> None:
in_progress = self._get_disagg_generation_trans_in_progress_requests()
for request in requests:
transfer_id = self._get_disagg_generation_transfer_id(request)
if transfer_id is not None:
in_progress[transfer_id] = request

def _queue_disagg_generation_trans_complete_draft_sync(
self, requests: Iterable[LlmRequest]) -> None:
pending = self._get_pending_disagg_generation_trans_complete_requests()
for request in requests:
if not getattr(request,
"is_disagg_generation_transmission_complete", False):
continue
transfer_id = self._get_disagg_generation_transfer_id(request)
if transfer_id is not None:
pending[transfer_id] = request

def _queue_disagg_generation_trans_complete_draft_sync_by_ids(
self, transfer_ids: Iterable[int]) -> None:
in_progress = self._get_disagg_generation_trans_in_progress_requests()
completed_requests = []
for transfer_id in transfer_ids:
request = in_progress.pop(transfer_id, None)
if request is not None:
completed_requests.append(request)
self._queue_disagg_generation_trans_complete_draft_sync(
completed_requests)

@staticmethod
def _resolve_disagg_generation_trans_complete_draft_tokens(
context_phase_draft_tokens: Optional[Iterable[int]],
enable_spec_decode: bool, max_total_draft_tokens: int) -> List[int]:
draft_tokens = ([] if context_phase_draft_tokens is None else
list(context_phase_draft_tokens))
if not draft_tokens and enable_spec_decode:
return [0] * max_total_draft_tokens
return draft_tokens

@staticmethod
def _sync_disagg_generation_trans_complete_draft_tokens(
requests: Iterable[LlmRequest],
enable_spec_decode: bool = False,
max_total_draft_tokens: int = 0) -> None:
for request in requests:
if not getattr(request,
"is_disagg_generation_transmission_complete", False):
continue

context_phase_params = request.context_phase_params
if context_phase_params is None:
continue

request.py_draft_tokens = PyExecutor._resolve_disagg_generation_trans_complete_draft_tokens(
context_phase_params.draft_tokens, enable_spec_decode,
max_total_draft_tokens)
request.draft_tokens = request.py_draft_tokens
request.py_draft_pages_allocated = len(request.py_draft_tokens)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

def _sync_pending_disagg_generation_trans_complete_draft_tokens(
self) -> None:
pending = self._get_pending_disagg_generation_trans_complete_requests()
if not pending:
return
max_total_draft_tokens = getattr(self.model_engine,
"max_total_draft_tokens",
self.max_total_draft_tokens)
self._sync_disagg_generation_trans_complete_draft_tokens(
pending.values(), self.model_engine.enable_spec_decode,
max_total_draft_tokens)
pending.clear()

@staticmethod
def _get_generation_num_draft_tokens(request: LlmRequest) -> int:
Comment thread
SimengLiu-nv marked this conversation as resolved.
py_draft_tokens = getattr(request, "py_draft_tokens", None)
if py_draft_tokens is None:
return request.num_draft_tokens
return max(len(py_draft_tokens), request.num_draft_tokens)

def _get_disagg_transfer_admission_controller(
self) -> DisaggTransferAdmissionController:
controller = getattr(self, "_disagg_transfer_admission_controller",
Expand Down Expand Up @@ -3649,6 +3769,9 @@ def _prepare_and_schedule_batch(self):
continue
request.draft_tokens = [0] * self.max_total_draft_tokens

if self.kv_cache_transceiver:
self._sync_pending_disagg_generation_trans_complete_draft_tokens()

scheduled_batch, scheduler_fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule(
)

Expand Down Expand Up @@ -5200,8 +5323,9 @@ def _compute_scheduled_tokens(context_requests, generation_requests):
else:
compute = max(1, remaining - reusable_in_chunk)
num_scheduled_ctx_tokens += compute
num_scheduled_gen_tokens = sum(1 + gen_req.num_draft_tokens
for gen_req in generation_requests)
num_scheduled_gen_tokens = sum(
1 + PyExecutor._get_generation_num_draft_tokens(gen_req)
for gen_req in generation_requests)
return num_scheduled_ctx_tokens + num_scheduled_gen_tokens

def _waiting_requests(self, context_requests: list[LlmRequest],
Expand Down Expand Up @@ -5839,7 +5963,10 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch):
ctx_draft_tokens = [
0
] * self.model_engine.max_total_draft_tokens
req.py_draft_tokens = [] if ctx_draft_tokens is None else ctx_draft_tokens
req.py_draft_tokens = [] if ctx_draft_tokens is None else list(
ctx_draft_tokens)
Comment thread
SimengLiu-nv marked this conversation as resolved.
req.draft_tokens = req.py_draft_tokens
req.py_draft_pages_allocated = len(req.py_draft_tokens)
beam_width = req.py_beam_width
if not self._update_sampler_state_for_disagg_gen_request(
req, beam_width, first_gen_tokens):
Expand Down Expand Up @@ -5981,6 +6108,8 @@ def _recv_disagg_gen_cache(self, new_gen_reqs):
if self._is_disagg_gen_only_no_context_benchmark():
for req in new_gen_reqs:
req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE
self._queue_disagg_generation_trans_complete_draft_sync(
new_gen_reqs)
return

if not self._uses_async_disagg_gen_transfer():
Expand All @@ -5991,10 +6120,13 @@ def _recv_disagg_gen_cache(self, new_gen_reqs):
self.kv_cache_transceiver.request_and_receive_sync(req)
if req.state == LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE:
self._sync_disagg_transfer_made_progress = True
self._queue_disagg_generation_trans_complete_draft_sync(
[req])
self._check_cache_transfer_errors("generation requests")
return

for req in new_gen_reqs:
self._track_disagg_generation_trans_in_progress([req])

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Entries added here are only popped via check_gen_transfer_status, but a user cancel of an in-flight gen transfer deletes the recv session synchronously (transceiver.cancel_request), so the rid never resurfaces and this entry (plus the LlmRequest it holds) leaks for the process lifetime. Could we pop it on the terminate/cancel path? (The timeout path may need the same.) Thanks!

self.kv_cache_transceiver.request_and_receive_async(req)

if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None:
Expand Down Expand Up @@ -6147,10 +6279,19 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0):
def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0):
result = self.kv_cache_transceiver.check_gen_transfer_status(atLeastNum)
if isinstance(result, tuple):
_, _, cancelled_reqs = result
completed_request_ids, failed_request_ids, cancelled_reqs = result
self._queue_disagg_generation_trans_complete_draft_sync_by_ids(
completed_request_ids)
in_progress = self._get_disagg_generation_trans_in_progress_requests(
)
for request_id in failed_request_ids:
in_progress.pop(request_id, None)
user_canceled_set = set(self.canceled_req_ids)
for req in cancelled_reqs:
req_id = req.py_request_id if not req.is_child else req.parent_request_id
transfer_id = self._get_disagg_generation_transfer_id(req)
if transfer_id is not None:
in_progress.pop(transfer_id, None)
if req_id not in user_canceled_set:
req.state = LlmRequestState.DISAGG_TRANS_ERROR
if not self._is_disagg_inflight_cancel_active():
Expand Down
9 changes: 4 additions & 5 deletions tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,11 +520,10 @@ def _init_llm(self, chat_template: Optional[str] = None):

# gpt-oss
self.harmony_adapter: HarmonyAdapter | None = None
disable_harmony = os.getenv("DISABLE_HARMONY_ADAPTER", "0") == "1"
if disable_harmony or self.model_config is None:
self.use_harmony = False
else:
self.use_harmony = (type(self.model_config).model_type == "gpt_oss")
from tensorrt_llm.tokenizer import uses_harmony_tokenization
model_type = (None if self.model_config is None else type(
self.model_config).model_type)
self.use_harmony = uses_harmony_tokenization(model_type)

self._ensure_post_processor_hook_supported(
self.use_harmony,
Expand Down
3 changes: 2 additions & 1 deletion tensorrt_llm/serve/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -798,14 +798,15 @@ def __init__(self,
tokens_per_block: Optional[int] = None,
custom_tokenizer: Optional[str] = None,
tokenizer_dir: Optional[str] = None,
use_harmony: Optional[bool] = None,
track_routed_blocks: bool = True,
load_weight: float = 0.25,
load_cap: float = float("inf"),
**kwargs):
super().__init__(server_role, servers, metadata_server_cfg,
metadata_server, **kwargs)
self._init_block_hashing(tokens_per_block, custom_tokenizer,
tokenizer_dir)
tokenizer_dir, use_harmony)
self._init_load_balancing(servers, use_tokens)
# TODO: use max_num_tokens? per server?
self._max_batch_size = max_batch_size
Expand Down
62 changes: 47 additions & 15 deletions tensorrt_llm/serve/router_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,19 +188,23 @@ 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,
) -> 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._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", use_harmony={self._use_harmony}"
)

def _get_tokenizer(self, model: str):
Expand All @@ -219,6 +223,41 @@ def _get_tokenizer(self, model: str):
self._tokenizers[model] = tokenizer.tokenizer
return self._tokenizers[model]

def _get_model_type(self, model: str) -> Optional[str]:
if model not in self._model_types:
model_path = self._tokenizer_dir or model
from tensorrt_llm.tokenizer import resolve_model_type_from_name_or_config

self._model_types[model] = resolve_model_type_from_name_or_config(model_path)
return self._model_types[model]

def _uses_harmony_tokenization(self, request: ChatCompletionRequest) -> bool:
from tensorrt_llm.tokenizer import uses_harmony_tokenization

return uses_harmony_tokenization(self._get_model_type(request.model), self._use_harmony)

@staticmethod
def _tool_dicts(request: ChatCompletionRequest) -> Optional[list[dict[str, object]]]:
if request.tools is None:
return None
return [
tool.model_dump() if hasattr(tool, "model_dump") else tool for tool in request.tools
]

def _tokenize_harmony_chat(self, request: ChatCompletionRequest) -> list[list[int]]:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we unify this implementation with openai_server.py:2053 added by @reasonsolo ?

from tensorrt_llm.serve import harmony_adapter

tools = self._tool_dicts(request) if request.tools else None
result = harmony_adapter.get_harmony_adapter().openai_to_harmony_tokens(
request.messages,
tools,
reasoning_effort=harmony_adapter.maybe_transform_reasoning_effort(
request.reasoning_effort
),
tool_choice=request.tool_choice,
)
return [result]

def _encode_with_prefix_cache(self, rendered: str, key: int, tokenizer) -> list[int]:
cache = getattr(self, "_tok_prefix_cache", None)
if cache is None:
Expand All @@ -241,26 +280,19 @@ def _tokenize(self, request: OpenAIRequest) -> list[list[int]]:
if isinstance(request, ChatCompletionRequest):
if request.prompt_token_ids is not None:
return [request.prompt_token_ids]
if self._uses_harmony_tokenization(request):
return self._tokenize_harmony_chat(request)
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 {}
)
chat_template_kwargs = dict(request.chat_template_kwargs or {})
chat_template_kwargs["tools"] = self._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,
tools=tool_dicts,
**chat_template_kwargs,
)
if isinstance(rendered, str):
Expand Down
4 changes: 4 additions & 0 deletions tensorrt_llm/tokenizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
load_custom_tokenizer,
load_hf_tokenizer,
maybe_fix_byte_level_tokenizer,
resolve_model_type_from_name_or_config,
tokenizer_factory,
uses_harmony_tokenization,
)

__all__ = [
Expand All @@ -21,7 +23,9 @@
"load_custom_tokenizer",
"load_hf_tokenizer",
"maybe_fix_byte_level_tokenizer",
"resolve_model_type_from_name_or_config",
"tokenizer_factory",
"uses_harmony_tokenization",
"_xgrammar_tokenizer_info",
"_llguidance_tokenizer_info",
]
Loading
Loading