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_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/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) 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() 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(