diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 4b85232c482c..c14d03d4433f 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -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 @@ -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( @@ -3355,6 +3361,120 @@ def _commit_kv_cache_stats(self, self.kv_cache_manager.commit_scheduled_kv_cache_stats( scheduled_batch) + @staticmethod + 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) + + 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: + 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", @@ -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( ) @@ -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], @@ -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) + 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): @@ -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(): @@ -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]) self.kv_cache_transceiver.request_and_receive_async(req) if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: @@ -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(): diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 3e1aa0d644db..89f45e1f1de1 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -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, diff --git a/tensorrt_llm/serve/router.py b/tensorrt_llm/serve/router.py index f0588a7e82f6..72f7c729c8bd 100644 --- a/tensorrt_llm/serve/router.py +++ b/tensorrt_llm/serve/router.py @@ -798,6 +798,7 @@ 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"), @@ -805,7 +806,7 @@ def __init__(self, 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 diff --git a/tensorrt_llm/serve/router_utils.py b/tensorrt_llm/serve/router_utils.py index a96d087acf06..33ee1d0658c1 100644 --- a/tensorrt_llm/serve/router_utils.py +++ b/tensorrt_llm/serve/router_utils.py @@ -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): @@ -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]]: + 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: @@ -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): diff --git a/tensorrt_llm/tokenizer/__init__.py b/tensorrt_llm/tokenizer/__init__.py index c6f7a8331ace..92c5b660c374 100644 --- a/tensorrt_llm/tokenizer/__init__.py +++ b/tensorrt_llm/tokenizer/__init__.py @@ -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__ = [ @@ -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", ] diff --git a/tensorrt_llm/tokenizer/tokenizer.py b/tensorrt_llm/tokenizer/tokenizer.py index 1422cc3804b5..478e2f872be7 100644 --- a/tensorrt_llm/tokenizer/tokenizer.py +++ b/tensorrt_llm/tokenizer/tokenizer.py @@ -117,6 +117,40 @@ def _load_tokenizer_config_inherits( return {k: tcfg[k] for k in _TOKENIZER_CONFIG_INHERIT_KEYS if k in tcfg} +def resolve_model_type_from_name_or_config(model: str) -> Optional[str]: + """Resolve the model type from a model alias or config.json when available.""" + normalized_model = model.lower().replace("_", "-") + if "gpt-oss" in normalized_model or "gptoss" in normalized_model: + return "gpt_oss" + + import json + config_path = os.path.join(model, "config.json") + if not os.path.isfile(config_path): + return None + + try: + with open(config_path, encoding="utf-8") as config_file: + config = json.load(config_file) + except (OSError, json.JSONDecodeError) as e: + logger.debug(f"Failed to read model config for {model}: {e}") + return None + + if not isinstance(config, dict): + return None + raw_model_type = config.get("model_type") + return raw_model_type if isinstance(raw_model_type, str) else None + + +def uses_harmony_tokenization(model_type: Optional[str], + use_harmony: Optional[bool] = None) -> bool: + """Return whether chat tokenization should use the GPT-OSS Harmony path.""" + if os.getenv("DISABLE_HARMONY_ADAPTER", "0") == "1": + return False + if use_harmony is not None: + return use_harmony + return model_type == "gpt_oss" + + def _fallback_to_fast_tokenizer(pretrained_model_dir: str, original_error: BaseException, **kwargs): """Bypass AutoTokenizer's HF-config path with PreTrainedTokenizerFast. diff --git a/tests/integration/defs/disaggregated/test_configs/disagg_config_2ctx_gptoss_eagle_trtllm.yaml b/tests/integration/defs/disaggregated/test_configs/disagg_config_2ctx_gptoss_eagle_trtllm.yaml new file mode 100644 index 000000000000..7abf3cd32c1d --- /dev/null +++ b/tests/integration/defs/disaggregated/test_configs/disagg_config_2ctx_gptoss_eagle_trtllm.yaml @@ -0,0 +1,67 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +model: gpt_oss/gpt-oss-120b +hostname: localhost +backend: pytorch +cuda_graph_config: null +speculative_config: + decoding_type: Eagle + max_draft_len: 3 + eagle3_one_model: true + speculative_model: gpt_oss/gpt-oss-120b-Eagle3 +context_servers: + num_instances: 2 + tensor_parallel_size: 1 + pipeline_parallel_size: 1 + moe_expert_parallel_size: 1 + enable_attention_dp: false + max_num_tokens: 20000 + max_seq_len: 131072 + max_batch_size: 16 + trust_remote_code: true + enable_chunked_prefill: true + disable_overlap_scheduler: true + router: + type: kv_cache_aware + kv_cache_config: + enable_block_reuse: true + enable_partial_reuse: true + event_buffer_max_size: 1024 + free_gpu_memory_fraction: 0.8 + moe_config: + backend: TRTLLM + print_iter_log: true + cache_transceiver_config: + backend: DEFAULT + max_tokens_in_buffer: 131072 +generation_servers: + num_instances: 1 + tensor_parallel_size: 1 + pipeline_parallel_size: 1 + moe_expert_parallel_size: 1 + enable_attention_dp: false + max_num_tokens: 512 + max_seq_len: 131072 + max_batch_size: 16 + trust_remote_code: true + enable_chunked_prefill: true + disable_overlap_scheduler: true + kv_cache_config: + enable_block_reuse: true + enable_partial_reuse: true + free_gpu_memory_fraction: 0.8 + moe_config: + backend: TRTLLM + cuda_graph_config: + enable_padding: true + batch_sizes: + - 1 + - 2 + - 4 + - 8 + - 16 + print_iter_log: true + cache_transceiver_config: + backend: DEFAULT + max_tokens_in_buffer: 131072 diff --git a/tests/integration/defs/disaggregated/test_disaggregated.py b/tests/integration/defs/disaggregated/test_disaggregated.py index 05c905db31bb..4e13c21b5480 100644 --- a/tests/integration/defs/disaggregated/test_disaggregated.py +++ b/tests/integration/defs/disaggregated/test_disaggregated.py @@ -339,6 +339,8 @@ def get_test_config(test_desc, example_dir, test_root): f"{test_configs_root}/disagg_config_ctxtp2_gentp2_gptoss_tllm.yaml", "gpt_oss_120b_eagle_triton_stress": f"{test_configs_root}/disagg_config_ctxtp2_gentp2_gptoss_eagle_triton.yaml", + "gpt_oss_120b_eagle_trtllm_2ctx": + f"{test_configs_root}/disagg_config_2ctx_gptoss_eagle_trtllm.yaml", "gpt_oss_120b_eagle_trtllm_stress": f"{test_configs_root}/disagg_config_ctxtp2_gentp2_gptoss_eagle_trtllm.yaml", "gpt_oss_120b_triton_stress": @@ -2706,10 +2708,20 @@ def test_disaggregated_deepseek_v3_lite_bf16_tllm_gen_helix( @skip_pre_blackwell @pytest.mark.skip_less_device(4) -@pytest.mark.parametrize("model_path", ['gpt_oss/gpt-oss-120b']) +@pytest.mark.parametrize(("model_path", "test_desc"), [ + pytest.param('gpt_oss/gpt-oss-120b', + 'gpt_oss_120b_harmony', + marks=pytest.mark.skip_less_device(4), + id='gpt_oss/gpt-oss-120b'), + pytest.param('gpt_oss/gpt-oss-120b', + 'gpt_oss_120b_eagle_trtllm_2ctx', + marks=pytest.mark.skip_less_device(8), + id='gpt_oss_120b_eagle_trtllm_2ctx'), +]) def test_disaggregated_gpt_oss_120b_harmony(disaggregated_test_root, disaggregated_example_root, - llm_venv, model_path): + llm_venv, model_path: str, + test_desc: str) -> None: model_dir = f"{llm_models_root()}/{model_path}" setup_model_symlink(llm_venv, model_dir, model_path) @@ -2719,8 +2731,10 @@ def test_disaggregated_gpt_oss_120b_harmony(disaggregated_test_root, env["TIKTOKEN_RS_CACHE_DIR"] = tiktoken_vocab env["TIKTOKEN_ENCODINGS_BASE"] = tiktoken_vocab + num_iters = 1 if test_desc == "gpt_oss_120b_eagle_trtllm_2ctx" else 5 run_disaggregated_test(disaggregated_example_root, - "gpt_oss_120b_harmony", + test_desc, + num_iters=num_iters, env=env, model_path=model_dir, cwd=llm_venv.get_working_directory()) diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index cc96800d204a..03cc380b701e 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -50,6 +50,7 @@ l0_dgx_b200: - disaggregated/test_disaggregated.py::test_disaggregated_deepseek_v3_lite_fp8_ucx[DeepSeek-V3-Lite-fp8] - disaggregated/test_disaggregated.py::test_disaggregated_deepseek_v3_lite_fp8_nixl[DeepSeek-V3-Lite-fp8] - disaggregated/test_disaggregated.py::test_disaggregated_gpt_oss_120b_harmony[gpt_oss/gpt-oss-120b] + - disaggregated/test_disaggregated.py::test_disaggregated_gpt_oss_120b_harmony[gpt_oss_120b_eagle_trtllm_2ctx] - accuracy/test_llm_api_pytorch.py::TestDeepSeekR1::test_nvfp4_multi_gpus[latency_adp_lmtp_tp4] - accuracy/test_llm_api_pytorch.py::TestMiniMaxM2::test_4gpus[attention_dp=False-cuda_graph=True-overlap_scheduler=True-tp_size=4-ep_size=4] TIMEOUT (60) - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_4gpus[pp4-mtp_nextn=0-attention_dp=False-cuda_graph=False-overlap_scheduler=False-torch_compile=False] TIMEOUT (60) diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index a0d77293b616..422e6a03efc3 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -417,10 +417,28 @@ def _make_ctx_request( return req -def _make_gen_request(num_draft_tokens=0): +def _make_gen_request(num_draft_tokens: int = 0) -> Mock: """Helper to create a mock generation request.""" req = Mock() req.num_draft_tokens = num_draft_tokens + req.py_draft_tokens = None + req.is_disagg_generation_transmission_complete = False + return req + + +def _make_disagg_trans_complete_request( + draft_tokens: list[int] | None, request_id: int = 0 +) -> Mock: + req = Mock() + req.py_request_id = request_id + req.request_id = request_id + req.py_disaggregated_params = None + req.is_disagg_generation_transmission_complete = True + req.context_phase_params = Mock(draft_tokens=draft_tokens) + req.py_draft_tokens = [] + req.draft_tokens = [] + req.py_draft_pages_allocated = 0 + req.num_draft_tokens = 0 return req @@ -811,6 +829,22 @@ def complete_or_error(req): charge_budget=False, ) + def test_gen_transfer_status_tracks_completed_requests_for_deferred_sync(self): + executor = object.__new__(PyExecutor) + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.check_gen_transfer_status.return_value = ([7], [], []) + trans_complete = _make_disagg_trans_complete_request([11, 12], request_id=7) + executor._disagg_generation_trans_in_progress_requests = {7: trans_complete} + executor._pending_disagg_generation_trans_complete_requests = {} + executor.canceled_req_ids = [] + executor._is_disagg_inflight_cancel_active = Mock(return_value=False) + executor._check_cache_transfer_errors = Mock() + + PyExecutor._check_disagg_gen_cache_transfer_status(executor, 0) + + assert executor._disagg_generation_trans_in_progress_requests == {} + assert executor._pending_disagg_generation_trans_complete_requests == {7: trans_complete} + def test_peer_cp_rank_enters_context_progress_poll(self): executor = object.__new__(PyExecutor) executor.dist = Mock(tp_size=1, cp_size=4, world_size=4) @@ -834,6 +868,33 @@ def test_peer_cp_rank_enters_context_progress_poll(self): @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") class TestDisaggTransferAdmissionPP: + def test_pp_schedule_syncs_completed_generation_transfer_before_scheduling(self): + executor = object.__new__(PyExecutor) + executor.dist = Mock( + rank=0, is_first_pp_rank=True, is_last_pp_rank=True, tp_size=1, cp_size=1 + ) + executor.enable_attention_dp = False + executor.kv_cache_transceiver = Mock() + executor.model_engine = Mock(enable_spec_decode=False, max_total_draft_tokens=0) + executor.max_total_draft_tokens = 0 + trans_complete = _make_disagg_trans_complete_request([11, 12], request_id=7) + executor._pending_disagg_generation_trans_complete_requests = {7: trans_complete} + scheduled_batch = ScheduledRequests() + + def schedule_after_sync(): + assert trans_complete.py_draft_tokens == [11, 12] + assert trans_complete.draft_tokens == [11, 12] + return scheduled_batch, [], 0 + + executor._schedule = Mock(side_effect=schedule_after_sync) + executor._apply_disagg_transfer_admission = Mock(return_value=([], False)) + + PyExecutor._pp_schedule_and_propagate(executor, microbatch_id=0) + + executor._schedule.assert_called_once_with() + executor._apply_disagg_transfer_admission.assert_called_once_with([]) + assert executor._pending_disagg_generation_trans_complete_requests == {} + def test_pp_schedule_applies_gate_before_serializing(self): executor = object.__new__(PyExecutor) executor.dist = Mock( @@ -975,6 +1036,48 @@ def test_generation_tokens(self): gen = [_make_gen_request(3), _make_gen_request(0)] assert PyExecutor._compute_scheduled_tokens([], gen) == (1 + 3) + (1 + 0) + def test_disagg_trans_complete_draft_tokens_are_scheduler_visible(self) -> None: + gen = [_make_gen_request(3) for _ in range(127)] + trans_complete = _make_disagg_trans_complete_request([11, 12, 13]) + gen.append(trans_complete) + + assert PyExecutor._compute_scheduled_tokens([], gen) == 127 * 4 + 1 + + PyExecutor._sync_disagg_generation_trans_complete_draft_tokens(gen) + + assert trans_complete.py_draft_tokens == [11, 12, 13] + assert trans_complete.draft_tokens == [11, 12, 13] + assert trans_complete.py_draft_pages_allocated == 3 + assert PyExecutor._compute_scheduled_tokens([], gen) == 128 * 4 + + def test_disagg_trans_complete_missing_draft_tokens_are_scheduler_visible(self) -> None: + trans_complete = _make_disagg_trans_complete_request(None) + PyExecutor._sync_disagg_generation_trans_complete_draft_tokens([trans_complete]) + + assert trans_complete.py_draft_tokens == [] + assert trans_complete.draft_tokens == [] + assert trans_complete.py_draft_pages_allocated == 0 + assert PyExecutor._compute_scheduled_tokens([], [trans_complete]) == 1 + + def test_disagg_trans_complete_missing_draft_tokens_use_spec_decode_budget(self) -> None: + trans_complete = _make_disagg_trans_complete_request(None) + PyExecutor._sync_disagg_generation_trans_complete_draft_tokens( + [trans_complete], enable_spec_decode=True, max_total_draft_tokens=3 + ) + + assert trans_complete.py_draft_tokens == [0, 0, 0] + assert trans_complete.draft_tokens == [0, 0, 0] + assert trans_complete.py_draft_pages_allocated == 3 + assert PyExecutor._compute_scheduled_tokens([], [trans_complete]) == 4 + + def test_sync_disagg_draft_tokens_ignores_regular_generation_requests(self) -> None: + gen = _make_gen_request(3) + + PyExecutor._sync_disagg_generation_trans_complete_draft_tokens([gen]) + + assert gen.py_draft_tokens is None + assert PyExecutor._compute_scheduled_tokens([], [gen]) == 4 + def test_mixed_context_and_generation(self): """Combined context (with chunk-shift) and generation tokens.""" # Non-last chunk: compute = 25 diff --git a/tests/unittest/disaggregated/test_router.py b/tests/unittest/disaggregated/test_router.py index ddaf12a0c8e2..cffb5ec9b0dd 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 @@ -179,6 +181,33 @@ def get_prompt_lengths(): return [100, 500, 10, 400, 2000, 100] +def _get_weather_tool() -> ChatCompletionToolsParam: + return ChatCompletionToolsParam(function=FunctionDefinition( + name="get_current_weather", + description="Get the current weather in a given location", + parameters={ + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "City and country", + }, + }, + "required": ["location"], + }, + )) + + +def _mock_tokenizer(token_ids: list[int] | None = None) -> mock.MagicMock: + if token_ids is None: + token_ids = [10, 20, 30] + tokenizer = mock.MagicMock() + tokenizer.apply_chat_template.return_value = token_ids + tokenizer.encode.return_value = token_ids + tokenizer.return_value = {"input_ids": token_ids} + return tokenizer + + @pytest.fixture def context_requests(): @@ -1758,6 +1787,383 @@ def test_create_router_conversation(): assert isinstance(router, ConversationRouter) +def test_tokenize_forwards_tools_and_chat_template_kwargs() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[11, 22, 33]) + documents = [{ + "title": "Weather policy", + "text": "Use Celsius for European weather." + }] + chat_template = "{{ messages[0]['content'] }}" + req = ChatCompletionRequest( + model="custom-chat-model", + messages=[{ + "role": "user", + "content": "what's the weather in Paris?" + }], + tools=[_get_weather_tool()], + documents=documents, + chat_template=chat_template, + chat_template_kwargs={"thinking": True}, + ) + + with mock.patch.object(router, "_get_tokenizer", return_value=tokenizer): + token_lists = router._tokenize(req) + + assert token_lists == [[11, 22, 33]] + tokenizer.apply_chat_template.assert_called_once() + assert req.prompt_token_ids == [11, 22, 33] + + kwargs = tokenizer.apply_chat_template.call_args.kwargs + tool_dicts = kwargs["tools"] + assert isinstance(tool_dicts, list) + assert tool_dicts[0]["function"]["name"] == "get_current_weather" + assert kwargs["documents"] == documents + assert kwargs["chat_template"] == chat_template + assert kwargs["thinking"] is True + + +@pytest.mark.parametrize(("tools", "expected_tools"), [(None, None), ([], [])]) +def test_tokenize_handles_absent_and_empty_tools( + tools: list[ChatCompletionToolsParam] | None, + expected_tools: list[ChatCompletionToolsParam] | None) -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[11, 22, 33]) + req = ChatCompletionRequest(model="custom-chat-model", + messages=[{ + "role": "user", + "content": "hello" + }], + tools=tools) + + with mock.patch.object(router, "_get_tokenizer", return_value=tokenizer): + router._tokenize(req) + + tokenizer.apply_chat_template.assert_called_once() + assert req.prompt_token_ids == [11, 22, 33] + kwargs = tokenizer.apply_chat_template.call_args.kwargs + assert kwargs["tools"] == expected_tools + + +def test_gpt_oss_tokenize_uses_harmony_tokens_for_router_hashes() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[900, 901, 902, 903]) + harmony_tokens = [100, 101, 102, 103, 104] + harmony = mock.MagicMock() + harmony.openai_to_harmony_tokens.return_value = harmony_tokens + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter", + return_value=harmony), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "maybe_transform_reasoning_effort", + return_value="medium"), mock.patch.object( + router, "_get_tokenizer", return_value=tokenizer): + req = ChatCompletionRequest( + model="openai/gpt-oss-20b", + messages=[{ + "role": "developer", + "content": "Use tools when useful." + }, { + "role": "user", + "content": "what's the weather in Paris?" + }], + tools=[_get_weather_tool()], + tool_choice="none", + reasoning_effort="medium", + ) + token_lists, block_hashes = router._tokenize_and_compute_block_hashes( + req) + + tokenizer.apply_chat_template.assert_not_called() + harmony.openai_to_harmony_tokens.assert_called_once() + assert token_lists == [harmony_tokens] + assert req.prompt_token_ids is None + + call_args = harmony.openai_to_harmony_tokens.call_args + assert call_args.args[0] == req.messages + tool_dicts = call_args.args[1] + assert isinstance(tool_dicts, list) + assert tool_dicts[0]["function"]["name"] == "get_current_weather" + assert call_args.kwargs["reasoning_effort"] == "medium" + assert call_args.kwargs["tool_choice"] == "none" + + expected_hashes = router._compute_block_hashes([harmony_tokens]) + assert block_hashes == expected_hashes + + +def test_gpt_oss_harmony_empty_tools_matches_chat_harmony_path() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + harmony_tokens = [100, 101, 102, 103, 104] + harmony = mock.MagicMock() + harmony.openai_to_harmony_tokens.return_value = harmony_tokens + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter", + return_value=harmony), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "maybe_transform_reasoning_effort", + return_value="medium"): + req = ChatCompletionRequest(model="openai/gpt-oss-20b", + messages=[{ + "role": "user", + "content": "hello" + }], + tools=[], + tool_choice="auto", + reasoning_effort="medium") + token_lists = router._tokenize(req) + + harmony.openai_to_harmony_tokens.assert_called_once() + assert token_lists == [harmony_tokens] + assert req.prompt_token_ids is None + + call_args = harmony.openai_to_harmony_tokens.call_args + assert call_args.args[0] == req.messages + assert call_args.args[1] is None + assert call_args.kwargs["reasoning_effort"] == "medium" + assert call_args.kwargs["tool_choice"] == "auto" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("tool_case", ["nonempty", "empty"]) +@pytest.mark.parametrize("template_case", + ["default", "documents", "custom_chat_template"]) +async def test_gpt_oss_router_and_server_create_same_tokenized_input( + tool_case: str, template_case: str) -> None: + 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) + request_tools = [_get_weather_tool()] if tool_case == "nonempty" else [] + request_kwargs: dict[str, object] = {} + if template_case == "documents": + request_kwargs["documents"] = [{ + "title": + "Weather policy", + "text": + "Use Celsius for European weather." + }] + elif template_case == "custom_chat_template": + request_kwargs["chat_template"] = "{{ messages[0]['content'] }}" + + def fake_harmony_tokens(messages: list[object], + tools: list[dict[str, object]] | None, + reasoning_effort: object | None = None, + tool_choice: object | None = None) -> list[int]: + tool_marker = 0 if tools is None else len(tools) + 10 + return [ + len(messages), + tool_marker, + len(str(reasoning_effort)), + len(str(tool_choice)), + ] + + harmony = mock.MagicMock() + harmony.openai_to_harmony_tokens.side_effect = fake_harmony_tokens + harmony.get_stop_tokens.return_value = [42] + + request = ChatCompletionRequest( + model="openai/gpt-oss-20b", + messages=[{ + "role": "developer", + "content": "Use tools when useful." + }, { + "role": "user", + "content": "what's the weather in Paris?" + }], + tools=request_tools, + tool_choice="auto", + reasoning_effort="medium", + stream=True, + max_completion_tokens=1, + **request_kwargs, + ) + router_request = copy.deepcopy(request) + server_request = copy.deepcopy(request) + + server = OpenAIServer.__new__(OpenAIServer) + server.harmony_adapter = harmony + server.await_disconnected = mock.AsyncMock() + server.allow_request_chat_template = template_case == "custom_chat_template" + server.model_config = SimpleNamespace(vocab_size=1000) + server.tokenizer = SimpleNamespace(tokenizer=SimpleNamespace( + vocab_size=1000)) + promise = mock.MagicMock() + promise.prompt_token_ids = [] + server.generator = SimpleNamespace( + args=SimpleNamespace(num_postprocess_workers=0), + generate_async=mock.MagicMock(return_value=promise), + ) + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter", + return_value=harmony), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "maybe_transform_reasoning_effort", + return_value="medium"), mock.patch( + "tensorrt_llm.serve.openai_server." + "maybe_transform_reasoning_effort", + return_value="medium"): + 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 is None + + router_call, server_call = harmony.openai_to_harmony_tokens.call_args_list + assert router_call.args == server_call.args + assert router_call.kwargs == server_call.kwargs + + +def test_use_harmony_flag_for_alias_model() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32, + use_harmony=True) + tokenizer = _mock_tokenizer(token_ids=[900, 901, 902, 903]) + harmony_tokens = [100, 101, 102, 103, 104] + harmony = mock.MagicMock() + harmony.openai_to_harmony_tokens.return_value = harmony_tokens + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter", + return_value=harmony), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "maybe_transform_reasoning_effort", + return_value="medium"), mock.patch.object( + router, "_get_tokenizer", return_value=tokenizer): + req = ChatCompletionRequest(model="served-model", + messages=[{ + "role": "user", + "content": "hello" + }]) + token_lists = router._tokenize(req) + + tokenizer.apply_chat_template.assert_not_called() + harmony.openai_to_harmony_tokens.assert_called_once() + assert token_lists == [harmony_tokens] + assert req.prompt_token_ids is None + + +def test_gpt_oss_config_model_type_uses_harmony(tmp_path: Path) -> None: + model_dir = tmp_path / "served-model" + model_dir.mkdir() + (model_dir / "config.json").write_text('{"model_type": "gpt_oss"}', + encoding="utf-8") + + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[900, 901, 902, 903]) + harmony_tokens = [100, 101, 102, 103, 104] + harmony = mock.MagicMock() + harmony.openai_to_harmony_tokens.return_value = harmony_tokens + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter", + return_value=harmony), mock.patch( + "tensorrt_llm.serve.harmony_adapter." + "maybe_transform_reasoning_effort", + return_value="medium"), mock.patch.object( + router, "_get_tokenizer", return_value=tokenizer): + req = ChatCompletionRequest(model=str(model_dir), + messages=[{ + "role": "user", + "content": "hello" + }]) + token_lists = router._tokenize(req) + + tokenizer.apply_chat_template.assert_not_called() + harmony.openai_to_harmony_tokens.assert_called_once() + assert token_lists == [harmony_tokens] + assert req.prompt_token_ids is None + + +def test_disable_harmony_adapter_uses_tokenizer_for_gpt_oss( + monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_HARMONY_ADAPTER", "1") + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[900, 901, 902, 903]) + + with mock.patch("tensorrt_llm.serve.harmony_adapter.get_harmony_adapter" + ) as get_harmony_adapter, mock.patch.object( + router, "_get_tokenizer", return_value=tokenizer): + req = ChatCompletionRequest(model="openai/gpt-oss-20b", + messages=[{ + "role": "user", + "content": "hello" + }]) + token_lists = router._tokenize(req) + + get_harmony_adapter.assert_not_called() + tokenizer.apply_chat_template.assert_called_once() + assert token_lists == [[900, 901, 902, 903]] + assert req.prompt_token_ids == [900, 901, 902, 903] + + +def test_tokenize_rewrites_completion_prompt_to_token_ids() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + tokenizer = _mock_tokenizer(token_ids=[10, 20, 30]) + + with mock.patch.object(router, "_get_tokenizer", return_value=tokenizer): + req = CompletionRequest(model="TinyLlama", prompt="hello") + token_lists = router._tokenize(req) + + assert token_lists == [[10, 20, 30]] + assert req.prompt == [10, 20, 30] + + +def test_tokenize_skipped_when_prompt_token_ids_already_set() -> None: + router = KvCacheAwareRouter(server_role=None, + servers=["server1"], + use_tokens=False, + max_batch_size=32, + tokens_per_block=32) + req = ChatCompletionRequest(model="openai/gpt-oss-20b", + messages=[{ + "role": "user", + "content": "hello" + }], + prompt_token_ids=[1, 2, 3]) + + with mock.patch.object(router, "_tokenize_harmony_chat") as harmony_path, \ + mock.patch.object(router, "_get_tokenizer") as get_tokenizer: + token_lists = router._tokenize(req) + + assert token_lists == [[1, 2, 3]] + harmony_path.assert_not_called() + get_tokenizer.assert_not_called() + + def test_block_hash_mixin_routes_through_transformers_tokenizer(): """``BlockHashMixin._get_tokenizer`` must call ``TransformersTokenizer.from_pretrained``. @@ -1850,20 +2256,24 @@ def _get_weather_tool() -> ChatCompletionToolsParam: )) -def _mock_tokenizer(token_ids=None): - """Return a mock tokenizer with a recorded apply_chat_template. +def _mock_tokenizer(token_ids: list[int] | None = None) -> mock.MagicMock: + """Return a mock tokenizer that emits the same token ids on each path. - ``apply_chat_template`` records its kwargs and returns the supplied - token id list. + ``apply_chat_template`` records its kwargs and ``encode`` / ``__call__`` + cover completion prompt tokenization. """ + if token_ids is None: + token_ids = [1, 2, 3, 4, 5] tok = mock.MagicMock() - tok.apply_chat_template.return_value = token_ids or [1, 2, 3, 4, 5] + tok.apply_chat_template.return_value = token_ids + tok.encode.return_value = token_ids + tok.return_value = {"input_ids": token_ids} return tok @pytest.mark.parametrize("router_class", [KvCacheAwareRouter, ConversationRouter]) -def test_tokenize_forwards_tools_and_chat_template_kwargs(router_class): +def test_tokenize_forwards_tools_and_kwargs_for_router_classes(router_class): """Regression test for PR #13232. ``BlockHashMixin._tokenize`` must forward the request's ``tools`` (as a @@ -1965,7 +2375,7 @@ def test_tokenize_preserves_empty_tools_list(): assert kwargs["tools"] == [] -def test_tokenize_skipped_when_prompt_token_ids_already_set(): +def test_tokenize_skips_tools_kwargs_when_prompt_token_ids_already_set(): """Skip tokenization when ``prompt_token_ids`` is already populated. When the caller pre-tokenizes (``prompt_token_ids`` set), the router