diff --git a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h index 94c46c3fb25f..97ba1b28c8de 100644 --- a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h +++ b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h @@ -222,6 +222,7 @@ class GenericLlmRequest mState = LlmRequestState::kENCODER_INIT; } + adoptContextPhaseDraftTokens(); initialize(*inputTokens, returnLogProbs, arrivalTime); } @@ -290,6 +291,7 @@ class GenericLlmRequest { mState = LlmRequestState::kENCODER_INIT; } + adoptContextPhaseDraftTokens(); initialize(inputTokens, returnLogProbs); } @@ -522,6 +524,7 @@ class GenericLlmRequest default: throw std::runtime_error("Unsupported request type found."); } + adoptContextPhaseDraftTokens(); initialize(req.getInputTokenIds(), req.getOutputConfig().returnLogProbs); } @@ -540,9 +543,29 @@ class GenericLlmRequest return mContextPhaseParams; } + /// @brief Get the number of generation tokens carried by context phase handoff. + /// @return Number of first generation tokens plus draft tokens. + [[nodiscard]] SizeType32 getNumContextPhaseGenerationTokens() const noexcept + { + if (!mContextPhaseParams.has_value()) + { + return 0; + } + + auto const& contextPhaseParams = mContextPhaseParams.value(); + auto numTokens = static_cast(contextPhaseParams.getFirstGenTokens().size()); + auto const& draftTokens = contextPhaseParams.getDraftTokens(); + if (draftTokens.has_value()) + { + numTokens += static_cast(draftTokens->size()); + } + return numTokens; + } + void setContextPhaseParams(executor::ContextPhaseParams contextPhaseParams) { mContextPhaseParams = std::move(contextPhaseParams); + adoptContextPhaseDraftTokens(); } /// @brief Get the state params of the context @@ -2253,6 +2276,20 @@ class GenericLlmRequest std::optional>> mAgentHierarchy{std::nullopt}; private: + void adoptContextPhaseDraftTokens() + { + if (hasDraftTokens() || !mContextPhaseParams.has_value()) + { + return; + } + + auto const& draftTokens = mContextPhaseParams.value().getDraftTokens(); + if (draftTokens.has_value() && !draftTokens->empty()) + { + mDraftTokens = std::make_shared(*draftTokens); + } + } + void initialize( VecTokens const& inputTokens, bool outputLogProbs, std::optional arrivalTime = std::nullopt) { diff --git a/cpp/tensorrt_llm/batch_manager/createNewDecoderRequests.cpp b/cpp/tensorrt_llm/batch_manager/createNewDecoderRequests.cpp index 04c3760be5c7..2d090a5612a5 100644 --- a/cpp/tensorrt_llm/batch_manager/createNewDecoderRequests.cpp +++ b/cpp/tensorrt_llm/batch_manager/createNewDecoderRequests.cpp @@ -69,10 +69,8 @@ void copySequenceLengths(RequestVector const& contextRequests, DecoderInputBuffe SizeType32 batchIdx{0}; for (auto const& llmReq : contextRequests) { - auto const disaggFirstGenTokenSize - = llmReq->getContextPhaseParams() ? llmReq->getContextPhaseParams().value().getFirstGenTokens().size() : 0; auto const currentSequenceLen - = llmReq->mPromptLen + llmReq->getMaxNumGeneratedTokens() + disaggFirstGenTokenSize; + = llmReq->mPromptLen + llmReq->getMaxNumGeneratedTokens() + llmReq->getNumContextPhaseGenerationTokens(); // Get position of the current sequence in the decoder auto const seqSlot = llmReq->mSeqSlot.value(); batchSlotsRange[batchIdx] = seqSlot; diff --git a/cpp/tests/unit_tests/batch_manager/llmRequestTest.cpp b/cpp/tests/unit_tests/batch_manager/llmRequestTest.cpp index 418178ff5023..4a897beec48d 100644 --- a/cpp/tests/unit_tests/batch_manager/llmRequestTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/llmRequestTest.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -790,6 +790,41 @@ TEST_F(LlmRequestTest, createResultDisaggContextComplete) EXPECT_TRUE(response->isSequenceFinal); } +TEST_F(LlmRequestTest, generationOnlyRequestAdoptsContextPhaseDraftTokens) +{ + VecTokens inputTokens{1, 2, 3, 4, 5}; + VecTokens firstGenTokens{100}; + VecTokens draftTokens{101, 102, 103}; + SizeType32 maxNewTokens{10}; + texec::IdType requestId{42}; + + texec::Request execReq(inputTokens, maxNewTokens); + execReq.setRequestType(texec::RequestType::REQUEST_TYPE_GENERATION_ONLY); + execReq.setContextPhaseParams( + texec::ContextPhaseParams{firstGenTokens, requestId, static_cast(nullptr), draftTokens}); + auto const expectedContextPhaseTokens = static_cast(firstGenTokens.size() + draftTokens.size()); + auto const expectedDraftTokens = static_cast(draftTokens.size()); + + tb::LlmRequest llmReq(requestId, execReq); + + EXPECT_TRUE(llmReq.isGenerationOnlyRequest()); + EXPECT_EQ(llmReq.getNumDraftTokens(), expectedDraftTokens); + EXPECT_EQ(*llmReq.getDraftTokens(), draftTokens); + EXPECT_EQ(llmReq.getNumContextPhaseGenerationTokens(), expectedContextPhaseTokens); + + texec::Request lateExecReq(inputTokens, maxNewTokens); + lateExecReq.setRequestType(texec::RequestType::REQUEST_TYPE_GENERATION_ONLY); + tb::LlmRequest lateLlmReq(requestId, lateExecReq); + EXPECT_EQ(lateLlmReq.getNumDraftTokens(), 0); + + lateLlmReq.setContextPhaseParams( + texec::ContextPhaseParams{firstGenTokens, requestId, static_cast(nullptr), draftTokens}); + + EXPECT_EQ(lateLlmReq.getNumDraftTokens(), expectedDraftTokens); + EXPECT_EQ(*lateLlmReq.getDraftTokens(), draftTokens); + EXPECT_EQ(lateLlmReq.getNumContextPhaseGenerationTokens(), expectedContextPhaseTokens); +} + INSTANTIATE_TEST_SUITE_P(LlmRequestTest, ParamTest, testing::Combine( // TODO: Support and add coverage for streamLLM diff --git a/tests/unittest/_torch/executor/test_request_utils.py b/tests/unittest/_torch/executor/test_request_utils.py index f42375c5e892..f2d6d6bc605b 100644 --- a/tests/unittest/_torch/executor/test_request_utils.py +++ b/tests/unittest/_torch/executor/test_request_utils.py @@ -1,3 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + """Tests for request_utils.py functions. This module tests: @@ -17,6 +20,7 @@ attach_py_objects_to_requests, can_process_attention_dp_request, derive_attention_dp_per_rank_request_cap, + executor_request_to_llm_request, get_from_waiting_queue, merge_helix_requests, merge_requests, @@ -95,6 +99,40 @@ def test_request_broadcaster_requires_conversation_params_attr(): RequestBroadcaster._collect_py_objects(None, source_items) +def test_executor_request_to_llm_request_adopts_context_phase_draft_tokens() -> None: + request_id = 42 + first_gen_tokens = [100] + draft_tokens = [101, 102, 103] + context_phase_params = trtllm.ContextPhaseParams( + first_gen_tokens, + request_id, + None, + draft_tokens, + None, + None, + ) + executor_request = trtllm.Request( + input_token_ids=[1, 2, 3], + max_tokens=10, + type=trtllm.RequestType.REQUEST_TYPE_GENERATION_ONLY, + context_phase_params=context_phase_params, + ) + + llm_request = executor_request_to_llm_request( + request_id, + executor_request, + child_req_ids=[], + exclude_last_generation_logits=False, + ) + + assert llm_request.is_generation_only_request() + assert llm_request.has_draft_tokens() + assert llm_request.num_draft_tokens == len(draft_tokens) + assert llm_request.draft_tokens == draft_tokens + assert llm_request.py_draft_tokens == draft_tokens + assert llm_request.context_phase_params.draft_tokens == draft_tokens + + def test_merge_helix_requests_with_padding(): """Test merge_helix_requests with basic valid input.""" diff --git a/tests/unittest/bindings/test_bindings_ut.py b/tests/unittest/bindings/test_bindings_ut.py index 210ea3378d3e..961748c25aad 100644 --- a/tests/unittest/bindings/test_bindings_ut.py +++ b/tests/unittest/bindings/test_bindings_ut.py @@ -1,3 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import json import pickle import tempfile @@ -417,6 +420,55 @@ def test_llm_request(): assert torch.equal(llm_request.draft_logits, logits) +def test_generation_only_llm_request_adopts_draft_tokens() -> None: + request_id = 42 + first_gen_tokens = [100] + draft_tokens = [101, 102, 103] + llm_request_type = _tb.internal.batch_manager.LlmRequestType + context_phase_params = _tb.executor.ContextPhaseParams( + first_gen_tokens, + request_id, + None, + draft_tokens, + None, + None, + ) + + llm_request = _tb.internal.batch_manager.LlmRequest( + request_id=request_id, + max_new_tokens=10, + sampling_config=_tb.SamplingConfig(1), + input_tokens=[1, 2, 3], + is_streaming=False, + llm_request_type=llm_request_type.LLMREQUEST_TYPE_GENERATION_ONLY, + context_phase_params=context_phase_params, + ) + + assert llm_request.is_generation_only_request + assert llm_request.has_draft_tokens() + assert llm_request.num_draft_tokens == len(draft_tokens) + assert llm_request.draft_tokens == draft_tokens + assert llm_request.context_phase_params.draft_tokens == draft_tokens + + late_llm_request = _tb.internal.batch_manager.LlmRequest( + request_id=request_id, + max_new_tokens=10, + sampling_config=_tb.SamplingConfig(1), + input_tokens=[1, 2, 3], + is_streaming=False, + llm_request_type=llm_request_type.LLMREQUEST_TYPE_GENERATION_ONLY, + ) + + assert late_llm_request.draft_tokens is None + assert late_llm_request.num_draft_tokens == 0 + + late_llm_request.context_phase_params = context_phase_params + + assert late_llm_request.has_draft_tokens() + assert late_llm_request.num_draft_tokens == len(draft_tokens) + assert late_llm_request.draft_tokens == draft_tokens + + def test_llm_request_kv_cache_transfer_metric_bindings(): request = _tb.internal.batch_manager.LlmRequest( request_id=0,