From ee0fd1baa6f569a30e4fb2e12fc674471cf98308 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Wed, 5 Aug 2026 22:43:36 -0700 Subject: [PATCH 1/3] [https://nvbugs/5807902][fix] Keep MiniMax-M3's separate draft KV cache in disaggregated serving The nvbug 5807902 WAR disables the separate draft KV cache manager whenever a cache transceiver is configured. MiniMax-M3 cannot take the shared-manager path that WAR forces: its cache manager declares supports_shared_draft_layers=False, and the drafter then inherits the target's tokens_per_block=128 pages, which miss the SM10x Eagle context cubins (the unfused-MHA fallback requests a 6.17 TiB workspace on a real 32K-token warmup) and hit the known tokens_per_block=128 trtllm-gen generation-kernel IMA. Both context and generation workers crashed during startup on every disaggregated Eagle3 attempt. Exempt MiniMax-M3 from the WAR so both worker roles keep the designed tokens_per_block=32 separate draft manager (symmetry is required for a consistent target pool layout across the KV transfer). Validated on Lyris GB300 (2xCTX TP2 + GEN TP4/ADP, NIXL): startup completes end to end, and the test_nvfp4_eagle3 chat-GSM8K acceptance workload measures AL 3.330 disagg vs 3.474 aggregated on the same build (drafter card reference 3.518). The remaining gap is the transceiver not transferring draft-layer KV, tracked separately. Signed-off-by: Zheyu Fu --- .../_torch/pyexecutor/py_executor_creator.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 93c8d8382a3f..9ee8e5a85172 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -514,8 +514,23 @@ def create_py_executor( # WAR for https://nvbugs/5807902 # Disable separate draft KV cache in disaggregated mode # Enable separate pool for None DI + Non-KVBM and Aggregated + KVBM + # + # MiniMax-M3 is exempt: its cache manager forbids sharing draft + # layers (supports_shared_draft_layers=False), so this WAR's + # shared-manager fallback does not exist for it. Both worker roles + # must keep the separate manager, or their target pool layouts + # diverge and disaggregated KV transfer breaks. Checked via the + # sparse-attention algorithm because the manager class is not yet + # resolved here ("minimax_m3" maps 1:1 to MiniMaxM3KVCacheManagerV2). if cache_transceiver_config is not None: - spec_config._allow_separate_draft_kv_cache = False + if is_minimax_m3(m3_sparse_config): + logger.warning( + "Disaggregated MiniMax-M3 keeps the separate draft KV " + "cache manager; draft-layer KV is not transferred, so " + "generation-side acceptance is reduced until the drafter " + "rebuilds its context window.") + else: + spec_config._allow_separate_draft_kv_cache = False # chunk_unit_size may be changed to 64 when using flash mla attn_runtime_features = AttentionRuntimeFeatures( From 5c18aaff0e8e971e32ff095f6eb60980f60040b8 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Thu, 6 Aug 2026 10:39:51 -0700 Subject: [PATCH 2/3] [TRTLLM-14019][feat] Transfer the separate draft KV cache in disaggregated serving One-model speculative decoding with a separate draft KV cache manager (e.g. MiniMax-M3's tokens_per_block=32 Eagle3 cache) previously lost the drafter's prompt KV at the CTX->GEN handoff: the transceiver only transferred the target manager's pools, so the generation-side drafter started every request against an unwarmed prompt window. Acceptance recovered only as generated tokens filled the drafter's attention window, which short prompts amortize (-4% AL on chat-GSM8K) but long-prompt regimes do not (-13% and worse as the prompt share of the window grows). Merge the draft manager's pools into the target's KVCachePageTable as additional attention layer groups, so NIXL registration, peer matching, the per-pool mappers, sessions and cancellation all ride the existing machinery: * AttentionLayerGroup gains an optional per-group tokens_per_block (draft 32 vs target 128 block math) carried through serialization. * merge_draft_page_table() appends draft groups with re-based pool indices and global layer ids offset by 1<<30; both peers apply the same constant, so the existing global-id-overlap + pool_role matching pairs draft groups with draft groups, and a peer without a draft manager degrades gracefully to today's behavior. * Head-mismatch mappers take per-layer-group KV head counts (the draft head count differs from the target's; rank-level remains the fallback so non-merged tables behave exactly as before). * KV slices source draft-group block ids from the draft manager and pin their cached-prefix to 0 (draft KV never participates in prefix reuse; the full prompt must transfer). * Target and draft pools register with the NIXL agent as separate batches: the V2 manager derives its VMM chunk size from the pool quota, so the two managers' chunk sizes differ, and the agent requires a uniform chunk size per registration batch (its region bookkeeping is already per-region). Validated on Lyris GB300 (2xCTX TP2 + GEN TP4/ADP, NIXL/Python transceiver, chat-GSM8K acceptance workload from test_nvfp4_eagle3): * acceptance length 3.330 -> 3.469 (aggregated baseline 3.474, drafter-card reference 3.518); acceptance rate 0.777 -> 0.823 (aggregated 0.825) * per-token-position probes show the early-step deficit healed (steps 1-3: 3.08 -> 3.54; steps 3-6: 2.75 -> 3.22 vs aggregated 3.88/3.69) and parity from step 6 on * 200/200 requests, zero failures; both worker roles log the draft cache registration Contains the M3 WAR exemption commit from #17341; will rebase once that merges. Signed-off-by: Zheyu Fu --- .../native/mixers/attention/peer.py | 42 ++++++---- .../_torch/disaggregation/native/peer.py | 11 +++ .../_torch/disaggregation/native/transfer.py | 81 ++++++++++++++++--- .../disaggregation/resource/kv_extractor.py | 52 ++++++++++++ .../_torch/disaggregation/resource/page.py | 8 ++ .../_torch/disaggregation/transceiver.py | 31 +++++-- tensorrt_llm/_torch/pyexecutor/_util.py | 8 +- .../_torch/pyexecutor/kv_cache_transceiver.py | 12 ++- 8 files changed, 213 insertions(+), 32 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py b/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py index 36aeb34ffdcd..c2ff3215a57d 100644 --- a/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py @@ -14,6 +14,7 @@ # limitations under the License. from collections.abc import Sequence +from typing import Optional import numpy as np @@ -177,9 +178,15 @@ def __init__( peer_bytes_per_layer: int, self_buffers_per_layer: int, peer_buffers_per_layer: int, + self_kv_heads: Optional[int] = None, + peer_kv_heads: Optional[int] = None, ): self._ri = self_ri self._peer_ri = peer_ri + # Head counts are per layer group (a merged draft group differs from + # the target's); None falls back to the rank-level value. + self_kv_heads = self_kv_heads or self_ri.attention.kv_heads_per_rank + peer_kv_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank self_tp_per_dp = self_ri.tp_size_per_dp_group peer_tp_per_dp = peer_ri.tp_size_per_dp_group @@ -212,28 +219,21 @@ def __init__( buffers_per_layer=peer_buffers_per_layer, side="peer", ) - bytes_per_head = self._bytes_per_head( - src_buffer_bytes, self._ri.attention.kv_heads_per_rank, side="local" - ) - peer_bytes_per_head = self._bytes_per_head( - dst_buffer_bytes, peer_ri.attention.kv_heads_per_rank, side="peer" - ) + bytes_per_head = self._bytes_per_head(src_buffer_bytes, self_kv_heads, side="local") + peer_bytes_per_head = self._bytes_per_head(dst_buffer_bytes, peer_kv_heads, side="peer") if bytes_per_head != peer_bytes_per_head: raise ValueError( f"HND bytes per head mismatch: local={bytes_per_head}, peer={peer_bytes_per_head}" ) - self._bytes_cont_heads = ( - min(self._ri.attention.kv_heads_per_rank, peer_ri.attention.kv_heads_per_rank) - * bytes_per_head - ) + self._bytes_cont_heads = min(self_kv_heads, peer_kv_heads) * bytes_per_head self._src_head_off, self._dst_head_off = self._compute_head_offsets( self_tp_per_dp, peer_tp_per_dp, self_tp_rank, peer_tp_rank, - self_kv_heads=self._ri.attention.kv_heads_per_rank, - peer_kv_heads=peer_ri.attention.kv_heads_per_rank, + self_kv_heads=self_kv_heads, + peer_kv_heads=peer_kv_heads, bytes_per_head=bytes_per_head, ) @@ -363,6 +363,8 @@ def __init__( peer_bytes_per_layer: int, self_buffers_per_layer: int, peer_buffers_per_layer: int, + self_kv_heads: Optional[int] = None, + peer_kv_heads: Optional[int] = None, ) -> None: # Deliberately do not call HNDHeadMismatchMapper.__init__: its offsets # assume HND-contiguous heads. Initialize the three attributes consumed @@ -375,8 +377,8 @@ def __init__( f"local={self_tpb}, peer={peer_tpb}" ) - self_heads = self_ri.attention.kv_heads_per_rank - peer_heads = peer_ri.attention.kv_heads_per_rank + self_heads = self_kv_heads or self_ri.attention.kv_heads_per_rank + peer_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank if self_buffers_per_layer != peer_buffers_per_layer: raise ValueError( "NHD buffer count per layer mismatch: " @@ -608,6 +610,8 @@ def build_kv_mapper( peer_bytes_per_layer: int, self_buffers_per_layer: int = 1, peer_buffers_per_layer: int = 1, + self_kv_heads: Optional[int] = None, + peer_kv_heads: Optional[int] = None, ) -> RegionMapperBase: """Pick the mapper for one view pair. @@ -633,7 +637,11 @@ def build_kv_mapper( peer_bytes_per_layer, ) - head_match, _ = self.head_match(peer_ri) + # Per-group head counts (merged draft groups differ from the target's + # rank-level value); None falls back to rank-level. + group_self_heads = self_kv_heads or self._ri.attention.kv_heads_per_rank + group_peer_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank + head_match = self._ri.attention.is_mla or group_self_heads == group_peer_heads if head_match: return IntactMapper( self_layer_offsets, @@ -653,6 +661,8 @@ def build_kv_mapper( peer_bytes_per_layer=peer_bytes_per_layer, self_buffers_per_layer=self_buffers_per_layer, peer_buffers_per_layer=peer_buffers_per_layer, + self_kv_heads=group_self_heads, + peer_kv_heads=group_peer_heads, ) return HNDHeadMismatchMapper( @@ -664,4 +674,6 @@ def build_kv_mapper( peer_bytes_per_layer=peer_bytes_per_layer, self_buffers_per_layer=self_buffers_per_layer, peer_buffers_per_layer=peer_buffers_per_layer, + self_kv_heads=group_self_heads, + peer_kv_heads=group_peer_heads, ) diff --git a/tensorrt_llm/_torch/disaggregation/native/peer.py b/tensorrt_llm/_torch/disaggregation/native/peer.py index ebea1dfba5ee..13959f42540e 100644 --- a/tensorrt_llm/_torch/disaggregation/native/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/peer.py @@ -342,6 +342,15 @@ def get_kv_map( pool_idx=peer_pi, ) + # Head counts are a per-layer-group property (a merged draft group's + # head count differs from the target's); fall back to the rank-level + # value for tables that predate per-group head metadata. + self_kv_heads = ( + getattr(self_lg, "kv_head_num_per_rank", 0) or self._ri.attention.kv_heads_per_rank + ) + peer_kv_heads = ( + getattr(peer_lg, "kv_head_num_per_rank", 0) or peer_ri.attention.kv_heads_per_rank + ) mapper = self._attention_policy.build_kv_mapper( peer_ri=peer_ri, mapper_kind=self_pv.mapper_kind, @@ -351,6 +360,8 @@ def get_kv_map( peer_bytes_per_layer=peer_bytes_per_layer, self_buffers_per_layer=self_buffers_per_layer, peer_buffers_per_layer=peer_buffers_per_layer, + self_kv_heads=self_kv_heads, + peer_kv_heads=peer_kv_heads, ) self._kv_map_cache[cache_key] = mapper diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 353444eeca6f..7d82dfd621a5 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -62,8 +62,12 @@ from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo from tensorrt_llm._torch.disaggregation.native.utils import get_local_ip from tensorrt_llm._torch.disaggregation.nixl.agent import NixlTransferAgent -from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 -from tensorrt_llm._torch.disaggregation.resource.page import MapperKind +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import ( + KVRegionExtractorV1, + build_page_table_from_manager, + merge_draft_page_table, +) +from tensorrt_llm._torch.disaggregation.resource.page import KVCachePageTable, MapperKind from tensorrt_llm._torch.disaggregation.resource.utils import get_unique_pool_memory_descs from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -809,9 +813,13 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write f"src/dst block count mismatch: {src_block_ids.size} vs " f"{dst_block_ids.size} (dst must not exceed src)" ) - tpb = extractor.page_table.tokens_per_block token_range = task._slice.token_range lg_info = extractor.page_table.layer_groups[self_lg] + # Merged draft groups carry their own page size; fall back to the + # table-wide value for regular target groups. + tpb = ( + getattr(lg_info, "tokens_per_block", None) or extractor.page_table.tokens_per_block + ) window_size = getattr(lg_info, "sliding_window_size", None) # Block lists are the suffix of [..., slice_end); cached prefix @@ -2242,6 +2250,10 @@ class TransferWorkerConfig: tx_timeout_s: Optional[float] = None rx_timeout_s: Optional[float] = None bounce: Optional["Config"] = None + # Separate one-model draft KV cache manager whose (small) prompt KV is + # transferred alongside the target's. Its layer groups are merged into + # the page table with offset global ids and a per-group tokens_per_block. + draft_kv_cache_manager: Optional[KVCacheManager] = None class TransferWorker: @@ -2257,9 +2269,35 @@ def __init__(self, config: TransferWorkerConfig): config.device_id, self._aux_buffer.meta if self._aux_buffer is not None else None, ) + assert self._rank_info.page_table is not None + self._num_target_layer_groups = len(self._rank_info.page_table.layer_groups) + self._num_target_pool_groups = len(self._rank_info.page_table.pool_groups) + if config.draft_kv_cache_manager is not None: + if kvm.mapping.pp_size != 1: + raise ValueError( + "Draft KV cache transfer is not supported with pipeline " + "parallelism yet (layer_num_per_pp bookkeeping covers " + "target layers only)." + ) + draft_pt = build_page_table_from_manager(config.draft_kv_cache_manager) + self._rank_info.page_table = merge_draft_page_table( + self._rank_info.page_table, draft_pt + ) + logger.info( + "Registered separate draft KV cache for disaggregated transfer: " + f"{len(draft_pt.layer_groups)} layer group(s), " + f"tokens_per_block={draft_pt.tokens_per_block} " + f"(target uses {self._rank_info.attention.tokens_per_block})" + ) self._setup_peer_infrastructure(kvm) self._setup_transfer_engine() + @property + def num_target_layer_groups(self) -> int: + """Layer groups in the page table that belong to the target manager; + indices >= this are merged draft-manager groups.""" + return self._num_target_layer_groups + def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp: list[int]): assert self._rank_info is not None self._rank_info.sender_endpoints = endpoints @@ -2300,7 +2338,9 @@ def sweep_stale_req_infos(self): def _setup_peer_infrastructure(self, kvm: KVCacheManager): self._rank_info_server = RankInfoServer(self._rank_info) if kvm.mapping.rank == 0 else None - self._kv_extractor = KVRegionExtractorV1(kvm) + # Build the extractor from the (possibly draft-merged) page table so + # extraction covers every registered pool, not just the target's. + self._kv_extractor = KVRegionExtractorV1(self._rank_info.page_table) self._peer_registrar = PeerRegistrar(self._rank_info, self._kv_extractor) def _setup_transfer_engine(self): @@ -2339,13 +2379,36 @@ def _setup_transfer_engine(self): def _register_kv_cache(self): assert self._rank_info.page_table is not None - memory_descs = get_unique_pool_memory_descs( - self._rank_info.page_table, self._rank_info.device_id - ) - if memory_descs: + pt = self._rank_info.page_table + n_target = self._num_target_layer_groups + # Pools from different managers can use different VMM chunk sizes + # (the V2 manager derives its cuMemCreate granularity from the pool + # quota, and a merged draft manager's quota is much smaller). The + # NIXL agent requires a uniform chunk size within one registration + # batch, while its region bookkeeping is per-region — so register + # the target's and the draft's pools as separate batches. + sub_tables = [ + KVCachePageTable(pt.tokens_per_block, pt.layer_groups[:n_target], pt.pool_groups) + ] + if n_target < len(pt.layer_groups): + sub_tables.append( + KVCachePageTable(pt.tokens_per_block, pt.layer_groups[n_target:], pt.pool_groups) + ) + for batch_idx, sub_table in enumerate(sub_tables): + memory_descs = get_unique_pool_memory_descs(sub_table, self._rank_info.device_id) + if not memory_descs: + continue + if batch_idx > 0: + # Keep descriptor names unique across registration batches. + memory_descs = [ + (ptr, size, dev, f"{name}_draft{batch_idx}") + for ptr, size, dev, name in memory_descs + ] reg_memory_desc = RegMemoryDescs("VRAM", memory_descs) self._agent.register_memory(reg_memory_desc) - logger.debug(f"Registered KV cache memory with transfer agent: {memory_descs}") + logger.debug( + f"Registered KV cache memory batch {batch_idx} with transfer agent: {memory_descs}" + ) self._registered_mem.append(reg_memory_desc) def _register_aux_buffer(self): diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index 9248bf506bce..f8e4e6082696 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -648,3 +648,55 @@ def build_page_table_from_manager(manager) -> KVCachePageTable: return _build_page_table_v2(manager) else: return build_page_table(manager) + + +# Offset applied to the draft manager's synthetic global layer ids when its +# layer groups are merged into the target's page table. Keeps them +# collision-free against every target id scheme (plain pp_layers ids and the +# virtual-layer encodings alike), and both peers apply the same constant, so +# peer matching by global-id overlap pairs draft groups with draft groups. +DRAFT_GLOBAL_LAYER_ID_OFFSET = 1 << 30 + + +def merge_draft_page_table( + target_pt: KVCachePageTable, + draft_pt: KVCachePageTable, +) -> KVCachePageTable: + """Append a separate draft manager's page table to the target's. + + The merged table lets the disaggregated transfer machinery treat draft + layers as additional attention layer groups: NIXL registration, peer + matching (by offset global ids + pool_role) and the per-pool mappers all + operate on the merged table. Draft groups carry their own + ``tokens_per_block`` (the draft manager's page size may differ from the + target's, e.g. MiniMax-M3's 32 vs 128). + """ + pool_base = len(target_pt.pool_groups) + merged_groups: List[LayerGroup] = list(target_pt.layer_groups) + for lg in draft_pt.layer_groups: + if not isinstance(lg, AttentionLayerGroup): + raise ValueError( + "merge_draft_page_table supports attention draft layer groups " + f"only, got {type(lg).__name__}" + ) + merged_groups.append( + AttentionLayerGroup( + pool_group_idx=lg.pool_group_idx + pool_base, + kv_head_num_per_rank=lg.kv_head_num_per_rank, + sliding_window_size=lg.sliding_window_size, + local_layers=[ + LocalLayer( + local_layer_id=ll.local_layer_id, + global_layer_id=ll.global_layer_id + DRAFT_GLOBAL_LAYER_ID_OFFSET, + ) + for ll in lg.local_layers + ], + pool_views=lg.pool_views, + tokens_per_block=draft_pt.tokens_per_block, + ) + ) + return KVCachePageTable( + tokens_per_block=target_pt.tokens_per_block, + layer_groups=merged_groups, + pool_groups=list(target_pt.pool_groups) + list(draft_pt.pool_groups), + ) diff --git a/tensorrt_llm/_torch/disaggregation/resource/page.py b/tensorrt_llm/_torch/disaggregation/resource/page.py index b15f21764e61..831894c7ff81 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/page.py +++ b/tensorrt_llm/_torch/disaggregation/resource/page.py @@ -254,6 +254,11 @@ class AttentionLayerGroup(LayerGroup): sliding_window_size: Optional[int] = None local_layers: List[LocalLayer] = field(default_factory=list) pool_views: List[PoolView] = field(default_factory=list) + # Page size override for groups whose manager uses a different + # tokens_per_block than the table default (e.g. a separate one-model + # draft KV cache manager merged into the target's table). None means + # "use KVCachePageTable.tokens_per_block". + tokens_per_block: Optional[int] = None def to_dict(self) -> dict: return { @@ -262,16 +267,19 @@ def to_dict(self) -> dict: "sliding_window_size": self.sliding_window_size, "local_layers": [ll.to_dict() for ll in self.local_layers], "pool_views": [pv.to_dict() for pv in self.pool_views], + "tokens_per_block": self.tokens_per_block, } @classmethod def from_dict(cls, data: dict) -> "AttentionLayerGroup": + tokens_per_block = data.get("tokens_per_block") return cls( pool_group_idx=int(data["pool_group_idx"]), kv_head_num_per_rank=int(data["kv_head_num_per_rank"]), sliding_window_size=data.get("sliding_window_size"), local_layers=[LocalLayer.from_dict(x) for x in data.get("local_layers", [])], pool_views=[PoolView.from_dict(pv) for pv in data.get("pool_views", [])], + tokens_per_block=int(tokens_per_block) if tokens_per_block is not None else None, ) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index a85fb1b85796..e6a0332374d1 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -66,6 +66,7 @@ def __init__( dist: Distributed, kv_cache_manager: KVCacheManager, cache_transceiver_config: CacheTransceiverConfig, + draft_kv_cache_manager: Optional[KVCacheManager] = None, ): self._dist: Distributed = dist self._kv_cache_manager = kv_cache_manager @@ -77,6 +78,14 @@ def __init__( ) self._check_compatible() self._reuse_adapter: CacheReuseAdapter = create_cache_reuse_adapter(kv_cache_manager) + # One-model speculative decoding with a separate draft manager: the + # drafter's prompt KV is transferred alongside the target's so the + # generation-side drafter does not start from an unwarmed window. + self._draft_reuse_adapter: Optional[CacheReuseAdapter] = ( + create_cache_reuse_adapter(draft_kv_cache_manager) + if draft_kv_cache_manager is not None + else None + ) self._device_id = torch.cuda.current_device() logger.info(f"device_id: {self._device_id} in KvCacheTransceiverV2") @@ -94,8 +103,10 @@ def __init__( rx_timeout_s=self.kv_transfer_timeout_ms / 1000.0, # Size 0 turns bounce off; the block-count gate is internal (tuned via env). bounce=bounce_config_from_size(cache_transceiver_config.kv_cache_bounce_size_mb), + draft_kv_cache_manager=draft_kv_cache_manager, ) ) + self._num_target_layer_groups = self._transfer_worker.num_target_layer_groups self._dp_rank = mapping.tp_rank if mapping.enable_attention_dp else 0 self._context_info_endpoint = self._broadcast_context_endpoint() self._init_sync_policy() @@ -177,16 +188,19 @@ def __exit__(self, _exc_type, _exc_val, _exc_tb): def _create_kv_slice(self, req: LlmRequest) -> KVSlice: adapter = self._reuse_adapter - tpb = adapter.tokens_per_block assert self._page_table is not None layer_groups = self._page_table.layer_groups + n_target = self._num_target_layer_groups is_gen_only = req.is_generation_only_request() + # Draft groups never participate in prefix reuse (their manager runs + # with reuse off and the transfer must cover the full prompt), so + # their cached count is pinned to 0. cached_per_lg = ( - adapter.get_cached_token_count_per_layer_group(req, layer_groups) + adapter.get_cached_token_count_per_layer_group(req, layer_groups[:n_target]) if is_gen_only - else [0] * len(layer_groups) - ) + else [0] * n_target + ) + [0] * (len(layer_groups) - n_target) token_range = None if req.prompt_len > 0: @@ -206,7 +220,14 @@ def _create_kv_slice(self, req: LlmRequest) -> KVSlice: if isinstance(lg, MambaLayerGroup): groups.append(np.array([], dtype=np.int64)) continue - block_ids = adapter.get_block_ids(req, idx, lg) + # Merged draft groups use the draft manager's block ids (indexed + # by the draft manager's own group numbering) and page size. + tpb = getattr(lg, "tokens_per_block", None) or adapter.tokens_per_block + if idx >= n_target: + assert self._draft_reuse_adapter is not None + block_ids = self._draft_reuse_adapter.get_block_ids(req, idx - n_target, lg) + else: + block_ids = adapter.get_block_ids(req, idx, lg) # Limit to prompt_len blocks, matching C++ cacheFormatter behavior. total_blocks = (req.prompt_len + tpb - 1) // tpb if block_ids.size > total_blocks: diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 1c3e7bc4149b..232878d718fb 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -2711,9 +2711,15 @@ def create_py_executor_instance( if isinstance(kv_cache_manager, BaseMambaCacheManager): mamba_cache_manager = kv_cache_manager + # A separate one-model draft KV cache manager (e.g. MiniMax-M3's + # tokens_per_block=32 Eagle3 cache) joins the transfer so the + # generation-side drafter receives the prompt KV computed during + # context-side prefill. + draft_kv_cache_manager = resources.get( + ResourceManagerType.DRAFT_KV_CACHE_MANAGER) kv_cache_transceiver = create_kv_cache_transceiver( mapping, dist, kv_cache_manager, attention_type, - cache_transceiver_config, mamba_cache_manager) + cache_transceiver_config, mamba_cache_manager, draft_kv_cache_manager) waiting_queue_policy = (scheduler_config.waiting_queue_policy if scheduler_config is not None else diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index 9b9ea2556e1e..a3366ed0b129 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -117,7 +117,8 @@ def create_kv_cache_transceiver( kv_cache_manager: KVCacheManager, attention_type: AttentionTypeCpp, cache_transceiver_config: CacheTransceiverConfig, - mamba_cache_manager: Optional[BaseMambaCacheManager] = None): + mamba_cache_manager: Optional[BaseMambaCacheManager] = None, + draft_kv_cache_manager: Optional[KVCacheManager] = None): if cache_transceiver_config is None or cache_transceiver_config.backend is None: logger.info("cache_transceiver is disabled") return None @@ -192,9 +193,16 @@ def create_kv_cache_transceiver( KvCacheTransceiverV2 logger.info("Using KvCacheTransceiverV2") return KvCacheTransceiverV2(mapping, dist, kv_cache_manager, - cache_transceiver_config) + cache_transceiver_config, + draft_kv_cache_manager) # Default: use C++ transceiver (transceiver_runtime is None or "CPP") + if draft_kv_cache_manager is not None: + logger.warning( + "Separate draft KV cache transfer is only supported by the " + "Python transceiver (transceiver_runtime='PYTHON'); the draft " + "manager's prompt KV will not be transferred, reducing " + "generation-side acceptance.") return BindKvCacheTransceiver(mapping, dist, kv_cache_manager, attention_type, cache_transceiver_config, mamba_cache_manager) From a5c7584da82c744bfb5c03bcfba57def16a6aa5c Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Fri, 7 Aug 2026 00:36:05 -0700 Subject: [PATCH 3/3] [TRTLLM-14019][feat] Let the separate draft KV cache join prefix reuse With a separate draft KV cache manager, the drafter's blocks never entered prefix reuse: the mirror created draft caches without a radix lookup and unconditionally stopped committing. On any multi-turn workload with block reuse enabled, the target reuses the previous turns' prefix while the drafter's KV for those positions does not exist; worse, the pages backing that range are recycled from the free list and typically hold a previous request's draft KV at scrambled positions, actively degrading acceptance (measured -0.2 to -0.35 AL on turns 2+ for both aggregated and disaggregated serving). Fix, in the draft manager mirror: - Defer reuse-eligible draft cache creation until after the target manager's prepare (the target runs last in the resource sweep and only then is its reuse boundary known); the executor calls the new prepare_deferred_draft_reuse() after each prepare_resources sweep. The draft radix lookup is clamped to the target boundary: a draft hit beyond it would expose radix-shared blocks to drafter writes. - Explicitly zero-fill the gap [draft_hit, target_boundary): the drafter only writes the target's context chunk and has no target hidden states for reused positions, so the gap is unrecoverable by recompute. Zero K/V behaves like masked attention and is benign, unlike recycled stale KV. - Commit draft blocks (at transfer start for disaggregated context, at the decode transition otherwise). With gaps zeroed this is always safe, and draft reuse coverage then grows turn over turn instead of deadlocking on an empty pool: in steady state the draft hit tracks the target boundary exactly. Validation (MiniMax-M3 NVFP4 + Eagle3 draft3, GB300): - Copy-task turn probe, turns 2-4: disagg 3.03-3.11 -> 3.59-3.65, agg 3.01-3.10 -> 3.63-3.82 (turn-1 parity in all arms). - Real AgentX traces (256K ctx, c32, disagg): AL 2.79 -> 3.04 on top of the draft-KV transfer; draft hits track the target boundary up to 93.7K reused tokens with zero gap fills after bootstrap. Signed-off-by: Zheyu Fu --- .../_torch/pyexecutor/kv_cache_manager_v2.py | 163 +++++++++++++++++- tensorrt_llm/_torch/pyexecutor/py_executor.py | 43 +++++ 2 files changed, 200 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index d7d1484b6fc9..009934dfa983 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1131,12 +1131,23 @@ def append_to_kv_heads_per_layer( self.enable_block_reuse = kv_cache_config.enable_block_reuse self.enable_partial_reuse = kv_cache_config.enable_partial_reuse self.disk_prefetch_num_reqs = kv_cache_config.disk_prefetch_num_reqs + # Draft managers get their own ConversationManager instance so their + # radix tree is keyed identically to the target's under the + # per-conversation policy. enable_conversation_manager = ( - self.enable_block_reuse - and self.block_reuse_policy == BlockReusePolicy.PER_CONVERSATION - and not self.is_draft + self.enable_block_reuse and self.block_reuse_policy == BlockReusePolicy.PER_CONVERSATION ) self.conversation_manager = ConversationManager() if enable_conversation_manager else None + # Draft caches whose reuse coverage matched the target's boundary at + # creation (no unwritten gap); only these may commit to the draft + # radix tree — committing a cache with a gap would key blocks the + # drafter never wrote with real tokens and poison future reuse. + self._draft_commit_ready: set[int] = set() + # Deferred creation is only armed once the executor's post-prepare + # hook has been observed: if any scheduling path lacks the hook, the + # mirror falls back to creation-without-reuse instead of leaving the + # forward pass without a draft cache. + self._deferred_draft_hook_seen = False # With pipeline parallelism, multiple microbatches can be in-flight # simultaneously, so we need slots for all concurrent sequences. @@ -2449,6 +2460,15 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): for req in scheduled_batch.context_requests: kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None: + if self._defer_draft_creation(req): + # First sight of a reuse-eligible context request. + # This mirror runs BEFORE the target manager's + # prepare (the target is move_to_end'd), so the + # target's reuse boundary (context_current_position) + # is not known yet. Creation moves to + # prepare_deferred_draft_reuse(), which the executor + # calls after the full prepare sweep. + continue kv_cache = self._create_kv_cache( req.py_request_id, req.lora_task_id, @@ -2480,6 +2500,11 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): raise RuntimeError( f"Missing draft KV cache for generation request {req.py_request_id}" ) + if req.py_request_id in self._draft_commit_ready: + # Prefill is complete once the request decodes: commit + # the draft prompt blocks for prefix reuse. + self.try_commit_blocks(req) + self._draft_commit_ready.discard(req.py_request_id) if not self._resume_and_restore(req.py_request_id, kv_cache): raise RuntimeError( f"Failed to resume draft KV cache for request {req.py_request_id}" @@ -2496,6 +2521,129 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): f"{req.py_request_id}: could not resize to {new_cap} tokens" ) + def _defer_draft_creation(self, req: LlmRequest) -> bool: + """Whether a draft cache creation must wait for the target's prepare. + + Draft prefix reuse needs the target's reuse boundary + (req.context_current_position), which the target manager only sets + during its own prepare — and the target runs LAST in the resource + manager sweep. Disagg generation-init requests are excluded: their + flow calls the draft manager after the target explicitly and never + runs the deferred hook. + """ + return ( + self._deferred_draft_hook_seen + and self.enable_block_reuse + and not req.is_dummy + and not req.is_disagg_generation_init_state + ) + + def prepare_deferred_draft_reuse(self, scheduled_batch: ScheduledRequests) -> None: + """Create draft caches deferred by _prepare_draft_resources. + + Runs after the full resource-manager prepare sweep, so the target's + reuse boundary is known. The draft radix lookup is clamped to that + boundary: the drafter only recomputes positions the target + recomputes, so a draft hit beyond it would leave radix-shared blocks + exposed to drafter writes. A draft hit BELOW the boundary leaves a + gap [draft_hit, target_boundary) the drafter never writes (it has no + target hidden states there); those pages are explicitly zero-filled. + Zero K/V behaves like masked attention (near-zero logits, zero value + contribution), which measures ~0.5 AL better than the recycled stale + KV the pages would otherwise hold. With the gap zeroed, committing + is always safe: pooled chains contain real drafter KV everywhere the + drafter ran, and benign zeros elsewhere, so reuse coverage grows + turn over turn instead of deadlocking on an empty pool. + """ + if not self.is_draft: + return + self._deferred_draft_hook_seen = True + if not self.enable_block_reuse: + return + # Read the target's reuse boundary OUTSIDE request_context: + # LlmRequest keeps dual context positions (target/draft) switched by + # use_draft_model, and request_context(True, ...) flips every request + # to the draft-side cursor, which is always 0 here. The boundary we + # clamp against is the TARGET's committed prefix. + target_boundaries = { + req.py_request_id: req.context_current_position + for req in scheduled_batch.context_requests + } + with request_context(True, scheduled_batch): + for req in scheduled_batch.context_requests: + if self.kv_cache_map.get(req.py_request_id) is not None: + continue + if not self._defer_draft_creation(req): + continue + if self.conversation_manager is not None: + self.conversation_manager.prepare_request(req) + all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX) + target_committed = min( + len(all_tokens) - 1, target_boundaries.get(req.py_request_id, 0) + ) + reuse_tokens = ( + self._augment_tokens_for_block_reuse(all_tokens, req, end=target_committed) + if target_committed > 0 + else None + ) + kv_cache = self._create_kv_cache( + req.py_request_id, + req.lora_task_id, + reuse_tokens, + cache_salt=req.cache_salt, + is_dummy=req.is_dummy, + ) + if kv_cache is None: + raise RuntimeError( + f"Failed to create draft KV cache for request {req.py_request_id}" + ) + draft_hit = kv_cache.num_committed_tokens + if not self._resume_and_restore(req.py_request_id, kv_cache): + raise RuntimeError( + f"Failed to resume draft KV cache for request {req.py_request_id}" + ) + draft_len = get_draft_token_length(req) + capacity = ( + req.context_current_position + + req.context_chunk_size + + draft_len + + self.num_extra_kv_tokens + ) + if not kv_cache.resize(capacity): + raise RuntimeError( + f"Draft KV cache context resize failed for request " + f"{req.py_request_id}: could not resize to {capacity} tokens" + ) + # The drafter only writes the target's context chunk + # [target_boundary, len); the gap between the draft reuse hit + # and that boundary would otherwise expose recycled stale KV. + self._zero_fill_draft_gap(req.py_request_id, draft_hit, target_committed) + self._draft_commit_ready.add(req.py_request_id) + + def _zero_fill_draft_gap(self, request_id: int, start_tok: int, end_tok: int) -> None: + """Zero the draft K/V pages for token positions [start_tok, end_tok). + + Handles unaligned bounds: partial blocks are zeroed only over the + gap's token slice, so a reused tail below start_tok and drafter + writes above end_tok are preserved. Must run after resize() so the + gap's pages are allocated. + """ + if end_tok <= start_tok: + return + tpb = self.tokens_per_block + for global_layer in self.layer_offsets: + buffers = self.get_buffers(global_layer) + if buffers is None: + continue + page_indices = self.get_batch_cache_indices([request_id], layer_idx=global_layer)[0] + for blk in range(start_tok // tpb, min(-(-end_tok // tpb), len(page_indices))): + page = page_indices[blk] + if page == BAD_PAGE_INDEX: + continue + t0 = max(start_tok - blk * tpb, 0) + t1 = min(end_tok - blk * tpb, tpb) + buffers[page, :, t0:t1].zero_() + def _augment_tokens_for_block_reuse( self, tokens: Sequence[int], req: LlmRequest, start: int = 0, end: int | None = None ) -> Sequence[TokenIdExt]: @@ -3135,9 +3283,11 @@ def release_resources( return requests def try_commit_blocks(self, request: LlmRequest) -> None: - should_block_reuse = ( - self.enable_block_reuse and not self.is_draft and not request.is_dummy_request - ) + # Draft caches may only commit when their reuse coverage matched the + # target's boundary at creation (no unwritten gap in the prefix); + # see prepare_deferred_draft_reuse. + draft_ok = not self.is_draft or request.py_request_id in self._draft_commit_ready + should_block_reuse = self.enable_block_reuse and draft_ok and not request.is_dummy_request if not should_block_reuse: return @@ -3177,6 +3327,7 @@ def release_index_slot(self, request_id: int) -> None: def free_resources(self, request: LlmRequest, pin_on_release: bool = False): if self.conversation_manager is not None: self.conversation_manager.finish_request(request) + self._draft_commit_ready.discard(request.py_request_id) self._allocated_draft_lens.pop(request.py_request_id, None) kv_cache = self.kv_cache_map.pop(request.py_request_id, None) if kv_cache is None: diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index f9dbfd7ce382..6e791185cf4a 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -454,6 +454,16 @@ def start_transfer(self, request: LlmRequest): else: block_id = None + # Commit the draft cache's prompt blocks alongside the target's + # so the drafter's prefix is reusable by later conversation + # turns (try_commit_blocks gates internally on reuse being + # enabled and on gap-free draft reuse coverage). + draft_kv_mgr = self.resource_manager.resource_managers.get( + ResourceManagerType.DRAFT_KV_CACHE_MANAGER) + if draft_kv_mgr is not None and hasattr(draft_kv_mgr, + "try_commit_blocks"): + draft_kv_mgr.try_commit_blocks(request) + self._requests_in_transfer[req_id] = request self._request_transfer_metadata[ req_id] = self.RequestTransferMetadata(block_id) @@ -2627,6 +2637,17 @@ def _executor_loop_pp(self): self.resource_manager.prepare_resources(scheduled_batch) + # Draft prefix reuse: the draft manager defers reuse- + # eligible cache creation until after the target's + # prepare (which sets the reuse boundary); create those + # caches now, before the forward pass. + draft_kv_mgr = self.resource_manager.resource_managers.get( + ResourceManagerType.DRAFT_KV_CACHE_MANAGER) + if draft_kv_mgr is not None and hasattr( + draft_kv_mgr, "prepare_deferred_draft_reuse"): + draft_kv_mgr.prepare_deferred_draft_reuse( + scheduled_batch) + # The generation requests that do not have batch_idx # need to be in front of the batch due to the assumptions # made in model_engine.py::_forward_step. This is only important @@ -4078,6 +4099,17 @@ def _executor_loop(self): self.resource_manager.prepare_resources(scheduled_batch) + # Draft prefix reuse: the draft manager defers reuse- + # eligible cache creation until after the target's + # prepare (which sets the reuse boundary); create those + # caches now, before the forward pass. + draft_kv_mgr = self.resource_manager.resource_managers.get( + ResourceManagerType.DRAFT_KV_CACHE_MANAGER) + if draft_kv_mgr is not None and hasattr( + draft_kv_mgr, "prepare_deferred_draft_reuse"): + draft_kv_mgr.prepare_deferred_draft_reuse( + scheduled_batch) + if self.kv_connector_manager: self.kv_connector_manager.handle_metadata() @@ -4555,6 +4587,17 @@ def _executor_loop_overlap(self): self.resource_manager.prepare_resources(scheduled_batch) + # Draft prefix reuse: the draft manager defers reuse- + # eligible cache creation until after the target's + # prepare (which sets the reuse boundary); create those + # caches now, before the forward pass. + draft_kv_mgr = self.resource_manager.resource_managers.get( + ResourceManagerType.DRAFT_KV_CACHE_MANAGER) + if draft_kv_mgr is not None and hasattr( + draft_kv_mgr, "prepare_deferred_draft_reuse"): + draft_kv_mgr.prepare_deferred_draft_reuse( + scheduled_batch) + if self.kv_connector_manager: self.kv_connector_manager.handle_metadata()