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
14 changes: 5 additions & 9 deletions tensorrt_llm/_torch/disaggregation/native/transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -757,18 +757,14 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write
src_block_ids = src_block_ids_per_groups[self_lg]
dst_block_ids = dst_block_ids_per_groups[peer_lg]

# Speculative decoding: generation may have one extra draft-token block.
# Both sides trim block lists to ceil(prompt_len / tpb) in
# _create_kv_slice, so dst must never exceed src. A smaller dst
# (generation prefix-cache reuse) is handled via dst_start below.
block_diff = dst_block_ids.size - src_block_ids.size
if block_diff == 1:
logger.debug(
f"Trimming 1 extra dst block for draft tokens: "
f"src={src_block_ids.size}, dst={dst_block_ids.size}"
)
dst_block_ids = dst_block_ids[:-1]
elif block_diff > 1:
if block_diff > 0:
raise ValueError(
f"src/dst block count mismatch: {src_block_ids.size} vs "
f"{dst_block_ids.size} (expected diff <= 1)"
f"{dst_block_ids.size} (dst must not exceed src)"
)
tpb = extractor.page_table.tokens_per_block
token_range = task._slice.token_range
Expand Down
18 changes: 10 additions & 8 deletions tensorrt_llm/_torch/disaggregation/transceiver.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,12 +182,16 @@ def _create_kv_slice(self, req: LlmRequest) -> KVSlice:

token_range = None
if req.prompt_len > 0:
# Align with KV cache allocation (prepare_disagg_gen_init /
# _get_context_bytes), which reserves prompt_len +
# num_extra_kv_tokens slots for speculative decoding methods
# (e.g. EAGLE3) that consume extra KV positions per request.
num_extra_kv_tokens = getattr(self._kv_cache_manager, "num_extra_kv_tokens", 0) or 0
token_range = TokenRange(start=0, end=req.prompt_len + num_extra_kv_tokens)
# end must match the trimmed block list below (ceil(prompt_len / tpb)
# blocks). num_extra_kv_tokens slots (speculative decoding) are not
# transferred. In the previously added support for ctx disabling
# speculative decoding while gen enables it, both sides currently
# use prompt_len as the transfer range, so the ranges stay
# consistent.
# TODO: the accuracy impact of not transferring num_extra_kv_tokens
# on MTP and other speculative decoding paths is currently unclear;
# revisit whether these extra KV slots need to be transferred.
token_range = TokenRange(start=0, end=req.prompt_len)

groups = []
for idx, lg in enumerate(layer_groups):
Expand All @@ -196,8 +200,6 @@ def _create_kv_slice(self, req: LlmRequest) -> KVSlice:
continue
block_ids = adapter.get_block_ids(req, idx, lg)
# Limit to prompt_len blocks, matching C++ cacheFormatter behavior.
# Extra blocks from num_extra_kv_tokens (speculative decoding) have
# uninitialized KV data and must not be transferred.
total_blocks = (req.prompt_len + tpb - 1) // tpb
if block_ids.size > total_blocks:
block_ids = block_ids[:total_blocks]
Expand Down
38 changes: 26 additions & 12 deletions tests/unittest/disaggregated/test_cache_reuse_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,8 +269,8 @@ def test_start_eq_end_rejected(self):


# ---------------------------------------------------------------------------
# _create_kv_slice: default TokenRange spans prompt_len + num_extra_kv_tokens
# so transferred KV matches what resize_context / _get_context_bytes allocate.
# _create_kv_slice: default TokenRange spans prompt_len, matching the
# trimmed block list actually transferred.
# ---------------------------------------------------------------------------


Expand Down Expand Up @@ -310,26 +310,40 @@ def _build_transceiver_for_kv_slice(num_extra_kv_tokens: int, prompt_len: int):


class TestCreateKvSliceTokenRange:
"""Default TokenRange built by _create_kv_slice must align with KV-cache allocation.
"""token_range.end must be prompt_len, matching the trimmed block list.

KV cache allocation in resize_context (V2) and prepare_resources (V1) reserves
prompt_len + num_extra_kv_tokens slots whenever speculative decoding (e.g.
EAGLE3, MTP) consumes extra KV positions per request. The transferred token
range must cover the same span, otherwise the receiver under-receives KV.
The sender reconstructs total_blocks from token_range.end (ceil(end / tpb)),
so end must stay at prompt_len -- not prompt_len + num_extra_kv_tokens --
to match the blocks actually transferred.
"""

def test_includes_num_extra_kv_tokens(self):
def test_excludes_num_extra_kv_tokens(self):
prompt_len = 17
num_extra_kv_tokens = 7
transceiver, req = _build_transceiver_for_kv_slice(num_extra_kv_tokens, prompt_len)

kv_slice = transceiver._create_kv_slice(req)

assert kv_slice.token_range is not None
assert (kv_slice.token_range.start, kv_slice.token_range.end) == (
0,
prompt_len + num_extra_kv_tokens,
)
assert (kv_slice.token_range.start, kv_slice.token_range.end) == (0, prompt_len)

def test_extra_tokens_do_not_cross_block_boundary(self):
# Reconstructed total_blocks (ceil(end / tpb)) must match the blocks sent.
prompt_len = 16
num_extra_kv_tokens = 7
transceiver, req = _build_transceiver_for_kv_slice(num_extra_kv_tokens, prompt_len)
tpb = transceiver._reuse_adapter.tokens_per_block

# Setup must actually exercise a boundary crossing: prompt_len ends on a
# block boundary and the extra tokens would otherwise add a block.
assert prompt_len % tpb == 0
assert (prompt_len + num_extra_kv_tokens + tpb - 1) // tpb == prompt_len // tpb + 1

kv_slice = transceiver._create_kv_slice(req)

end = kv_slice.token_range.end
transferred_blocks = kv_slice.block_ids_per_layer_groups[0].size
assert (end + tpb - 1) // tpb == transferred_blocks

def test_defaults_to_prompt_len_when_no_extra(self):
prompt_len = 17
Expand Down
Loading