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
37 changes: 37 additions & 0 deletions cpp/include/tensorrt_llm/batch_manager/llmRequest.h
Original file line number Diff line number Diff line change
Expand Up @@ -222,6 +222,7 @@ class GenericLlmRequest
mState = LlmRequestState::kENCODER_INIT;
}

adoptContextPhaseDraftTokens();
initialize(*inputTokens, returnLogProbs, arrivalTime);
}

Expand Down Expand Up @@ -290,6 +291,7 @@ class GenericLlmRequest
{
mState = LlmRequestState::kENCODER_INIT;
}
adoptContextPhaseDraftTokens();
initialize(inputTokens, returnLogProbs);
}

Expand Down Expand Up @@ -522,6 +524,7 @@ class GenericLlmRequest
default: throw std::runtime_error("Unsupported request type found.");
}

adoptContextPhaseDraftTokens();
initialize(req.getInputTokenIds(), req.getOutputConfig().returnLogProbs);
}

Expand All @@ -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<SizeType32>(contextPhaseParams.getFirstGenTokens().size());
auto const& draftTokens = contextPhaseParams.getDraftTokens();
if (draftTokens.has_value())
{
numTokens += static_cast<SizeType32>(draftTokens->size());
}
return numTokens;
}

void setContextPhaseParams(executor::ContextPhaseParams contextPhaseParams)
{
mContextPhaseParams = std::move(contextPhaseParams);
adoptContextPhaseDraftTokens();
}

/// @brief Get the state params of the context
Expand Down Expand Up @@ -2253,6 +2276,20 @@ class GenericLlmRequest
std::optional<std::vector<std::tuple<std::string, int>>> 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<VecTokens>(*draftTokens);
}
}

void initialize(
VecTokens const& inputTokens, bool outputLogProbs, std::optional<TimePoint> arrivalTime = std::nullopt)
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
37 changes: 36 additions & 1 deletion cpp/tests/unit_tests/batch_manager/llmRequestTest.cpp
Original file line number Diff line number Diff line change
@@ -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");
Expand Down Expand Up @@ -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<void*>(nullptr), draftTokens});
auto const expectedContextPhaseTokens = static_cast<SizeType32>(firstGenTokens.size() + draftTokens.size());
auto const expectedDraftTokens = static_cast<SizeType32>(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<void*>(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
Expand Down
38 changes: 38 additions & 0 deletions tests/unittest/_torch/executor/test_request_utils.py
Original file line number Diff line number Diff line change
@@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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."""

Expand Down
52 changes: 52 additions & 0 deletions tests/unittest/bindings/test_bindings_ut.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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,
Expand Down
Loading