diff --git a/tensorrt_llm/_torch/attention_backend/fmha/flashinfer_trtllm_gen.py b/tensorrt_llm/_torch/attention_backend/fmha/flashinfer_trtllm_gen.py index ff2002c36d57..1e028bcb9c0c 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/flashinfer_trtllm_gen.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/flashinfer_trtllm_gen.py @@ -12,7 +12,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - """ FlashInfer TRTLLM-Gen FMHA @@ -66,7 +65,6 @@ TrtllmAttentionMetadata, ) - _MULTI_CTAS_KV_COUNTER_ALIGNMENT = 8 @@ -752,8 +750,15 @@ def _is_supported_with_reason( return False, f"non-positive tokens_per_block ({tokens_per_block})." if tokens_per_block & (tokens_per_block - 1) != 0: return False, f"tokens_per_block ({tokens_per_block}) that is not a power of 2." - if tokens_per_block not in self.SUPPORTED_TOKENS_PER_BLOCK: - supported = sorted(self.SUPPORTED_TOKENS_PER_BLOCK) + # P128 is not exported for every TRTLLM-Gen shape family, so keep it + # out of the global allowlist. Cache-manager views whose exact shapes + # are exported may opt in explicitly. + extra_tokens_per_block = getattr( + meta.kv_cache_manager, "trtllm_gen_extra_tokens_per_block", () + ) + supported_tokens_per_block = self.SUPPORTED_TOKENS_PER_BLOCK | set(extra_tokens_per_block) + if tokens_per_block not in supported_tokens_per_block: + supported = sorted(supported_tokens_per_block) return False, f"tokens_per_block ({tokens_per_block}). Supported: {supported}." return True, "" diff --git a/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py b/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py index c8d602c9f852..4cece7df218c 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py @@ -235,22 +235,6 @@ def run_msa_paged_gqa( ) return - if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx): - from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.trtllm_gen_dense_decode import ( - minimax_m3_trtllm_gen_dense_attention, - ) - - minimax_m3_trtllm_gen_dense_attention( - q_view, - kv_cache_manager, - layer_idx, - metadata, - sm_scale=sm_scale, - output=out_view, - kv_scale_quant_orig=kv_scale_quant_orig, - ) - return - k_paged, v_paged = msa_paged_kv(kv_cache_manager, layer_idx) # Leading query tokens fmha_sm100 must still run: the whole batch until a diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/cache_manager.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/cache_manager.py index 6424279988f9..4fe6db2fe223 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/cache_manager.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/cache_manager.py @@ -25,10 +25,8 @@ from __future__ import annotations -import os from typing import List, Optional, Sequence, Tuple -import numpy as np import torch from tensorrt_llm._torch.disaggregation.resource.page import MapperKind @@ -144,16 +142,9 @@ def derive_shared_draft_layout( ) -> tuple[list[int], Optional[int]]: """Locate the appended one-model draft tail in the manager's layer range. - ``num_layers`` is ambiguous at the creation site: for M3 + Eagle3 it - carries the pretrained TARGET count (60) while the per-layer - ``num_kv_heads`` list is already extended with the draft entries - (61); other flows pass the extended count directly. The heads list's - length is the unambiguous total, so anchor on it and fall back to - ``num_layers`` for scalar heads. - - Returns ``(draft_layer_ids, num_target_layers)``; the target range is - ``[0, num_target_layers)`` and the draft tail sits directly above it. - ``num_target_layers`` is ``None`` when neither input pins the range. + ``num_layers`` can describe either the pretrained target or the extended + target-plus-draft stack. The per-layer KV-head list is the unambiguous + total when available, while ``num_draft`` identifies the appended tail. """ total = ( len(num_kv_heads) @@ -165,8 +156,7 @@ def derive_shared_draft_layout( if num_layers is not None: total = max(total, int(num_layers)) num_target = total - max(0, int(num_draft)) - draft_ids = list(range(num_target, total)) - return draft_ids, num_target + return list(range(num_target, total)), num_target class MiniMaxM3KVCacheManagerV2(KVCacheManagerV2): @@ -188,36 +178,10 @@ class MiniMaxM3KVCacheManagerV2(KVCacheManagerV2): * ``disable_index_value_layer_ids`` — subset whose index-V is omitted. * ``sparse_index_dim`` — width of the index-K/V vectors. - * ``num_one_model_draft_layers`` — how many one-model draft layers - the creation site appended after the target's (0 when the drafter - is separate or speculation is off). + * ``num_one_model_draft_layers`` — number of appended shared draft + layers (0 for a separate drafter or when speculation is disabled). """ - # One-model speculative draft layers share this manager (unified KV - # cache): reuse, eviction, and disaggregated transfer then cover the - # drafter's KV natively. Its attention addresses the physical P32 buffer - # expansion through ``get_draft_subpage_view``. - supports_shared_draft_layers = True - - # WAR: the Eagle draft kernels break at tokens_per_block=128 (the MSA - # target's page size) — the SM103 context cubin is missing (its unfused - # fallback demands a multi-TiB workspace) and the generation kernel hits - # an illegal memory access — so the drafter runs at 32-token pages. This - # value sizes both the separate draft manager and the view's sub-pages. - # Retirement, once the kernels are fixed (WAR sites point here): - # 1. Validate with TRTLLM_M3_DRAFT_KV_TOKENS_PER_BLOCK=128; the view - # degenerates to the identity expansion (unit-tested). - # 2. Delete the WAR surface — MiniMaxM3DraftSubpageView, - # ``get_draft_subpage_view``, the ``add_dummy_requests`` override, - # and this attribute; the drafter then attends the shared manager - # directly (validated at acceptance parity, PR #17457). - draft_manager_tokens_per_block = 32 - # The separately allocated Eagle layer cannot inherit an NVFP4 target - # cache: the shipped TRTLLM-Gen set has no matching M3 P32 NVFP4 decode - # cubin. Keep the draft on its established FP8/P32 representation while - # sparse target layers use NVFP4/P128. - draft_manager_kv_cache_dtype = "fp8" - nvfp4_dense_tokens_per_block = 32 _main_kv_layout = "NHD" def __init__( @@ -263,9 +227,6 @@ def __init__( if sparse_index_dim is None: sparse_index_dim = int(getattr(sparse_attn_config, "sparse_index_dim", 0) or 0) or 128 - # One-model speculative decoding with shared draft layers appends the - # drafter's layers after the target's (dense, no MSA index cache); - # ``_create_kv_cache_manager`` passes the appended count explicitly. self._shared_draft_layer_ids, num_target_layers = derive_shared_draft_layout( num_layers, kwargs.get("num_kv_heads"), num_one_model_draft_layers ) @@ -309,10 +270,10 @@ def __init__( "[m3-kv] hybrid cache active: " f"{len(self.sparse_layer_ids)} sparse target layer(s)=NVFP4/P128, " f"{len(dense_target_layers)} dense target layer(s)=FP8/P128, " - f"{len(self._shared_draft_layer_ids)} shared Eagle layer(s)=FP8/P32" + f"{len(self._shared_draft_layer_ids)} shared Eagle layer(s)=FP8/P128" ) - self._draft_subpage_view_obj: Optional["MiniMaxM3DraftSubpageView"] = None + self._draft_kv_cache_view_obj: Optional["MiniMaxM3DraftKVCacheView"] = None if self._shared_draft_layer_ids and self.sparse_layer_ids and not self.is_draft: # Paired with the "view active" log at first dispatch. logger.info( @@ -353,18 +314,11 @@ def _build_cache_config(self, config): M3's 57 sparse target layers have a native MSA NVFP4 consumer. The three dense target layers and the appended one-model Eagle layer do not have a matching TRTLLM-Gen NVFP4 cubin, so those buffers retain - the proven FP8 representation. Target dense layers retain their - established P128 layout; only the shared Eagle layer uses physical - P32 pages because its SM100/SM103 kernels require that geometry. + the proven FP8/P128 representation. """ if self.dtype != DataType.NVFP4: return super()._build_cache_config(config) - physical_page = self.nvfp4_dense_tokens_per_block - assert config.tokens_per_block % physical_page == 0, ( - f"M3 logical page P{config.tokens_per_block} must be divisible by " - f"the dense/Eagle physical page P{physical_page}." - ) scale_roles = {Role.KEY_BLOCK_SCALE, Role.VALUE_BLOCK_SCALE} for layer in config.layers: local_layer_idx = int(layer.layer_id) @@ -374,14 +328,6 @@ def _build_cache_config(self, config): layer.buffers[:] = [ buffer for buffer in layer.buffers if buffer.role not in scale_roles ] - for buffer in layer.buffers: - if buffer.role not in (Role.KEY, Role.VALUE): - continue - if global_layer_idx in self._shared_draft_layer_ids: - buffer.size = ( - self.get_layer_bytes_per_token(local_layer_idx, buffer.role) * physical_page - ) - buffer.tokens_per_block_override = physical_page return super()._build_cache_config(config) def get_layer_bytes_per_token(self, local_layer_idx: int, data_role: Role): @@ -414,10 +360,6 @@ def is_fp8_dense_layer(self, layer_idx: int) -> bool: """Whether a dense target/Eagle layer is the FP8 half of hybrid KV.""" return self.dtype == DataType.NVFP4 and int(layer_idx) not in self.sparse_layer_ids - def is_fp8_subpaged_layer(self, layer_idx: int) -> bool: - """Whether the shared Eagle layer uses physical P32 FP8 pages.""" - return self.is_fp8_dense_layer(layer_idx) and int(layer_idx) in self._shared_draft_layer_ids - @property def uses_hybrid_nvfp4_kv_cache(self) -> bool: return self.dtype == DataType.NVFP4 @@ -428,8 +370,8 @@ def _build_pool_mapping_tensors(self): The base NVFP4 implementation assumes every layer has a scale buffer. Hybrid M3 deliberately omits scales from dense/Eagle layers, while the target metadata still needs the nested NVFP4 pointer envelope for the - sparse pools. Dense consumers use the direct P32 views below and the - shared Eagle view publishes itself as an ordinary FP8 manager. + sparse pools. Dense consumers use direct P128 views, and the shared + Eagle view publishes itself as an ordinary FP8 manager. """ if self.dtype != DataType.NVFP4: return super()._build_pool_mapping_tensors() @@ -501,48 +443,41 @@ def _build_pool_mapping_tensors(self): torch.tensor(mapping_rows, dtype=torch.int32, pin_memory=prefer_pinned()), ) - def get_draft_subpage_view(self) -> Optional["MiniMaxM3DraftSubpageView"]: - """Sub-page view over the shared drafter pool, or None. + def get_draft_kv_cache_view(self) -> Optional["MiniMaxM3DraftKVCacheView"]: + """Return a P128 view rooted at the shared draft layer's K page. Only meaningful on a target manager carrying appended one-model draft layers; built lazily so the manager's page tables exist. A method rather than a property so ``getattr`` fetches it without - executing it (see ``resolve_draft_kv_cache_manager``). + executing it (see ``get_draft_kv_cache_manager``). - Retires with the P128 Eagle kernel fixes; see - ``draft_manager_tokens_per_block``. + The view is required because M3's sparse layers add an index-K page to + the otherwise dense K/V mega-slot layout. """ if self.is_draft or not self._shared_draft_layer_ids: return None - if self._draft_subpage_view_obj is None: - subpage_tokens = ( - int(os.environ.get("TRTLLM_M3_DRAFT_KV_TOKENS_PER_BLOCK", 0) or 0) - or self.draft_manager_tokens_per_block - ) - self._draft_subpage_view_obj = MiniMaxM3DraftSubpageView( - self, - self._shared_draft_layer_ids, - subpage_tokens, + if self._draft_kv_cache_view_obj is None: + self._draft_kv_cache_view_obj = MiniMaxM3DraftKVCacheView( + self, self._shared_draft_layer_ids ) logger.info( - f"[unified-kv] draft sub-page view active " - f"(tokens_per_block={self._draft_subpage_view_obj.tokens_per_block}, " - f"flat_page_bound={self._draft_subpage_view_obj.blocks_in_primary_pool})" + "[unified-kv] native P128 draft view active " + f"(flat_page_bound={self._draft_kv_cache_view_obj.blocks_in_primary_pool})" ) - return self._draft_subpage_view_obj + return self._draft_kv_cache_view_obj def add_dummy_requests(self, *args, **kwargs): - """Drop the draft sub-page view before delegating. + """Drop the draft view before delegating. The base method mirrors dummy KV caches into a *separate* draft manager. With shared draft layers a dummy request's blocks already span the drafter's pool (pools allocate in lockstep per logical block), and the view owns no block lifecycle. - Retires with the P128 Eagle kernel fixes; see - ``draft_manager_tokens_per_block``. + This override remains necessary for the rooted P128 view because that + view owns no blocks and must not receive duplicate dummy allocations. """ - if isinstance(kwargs.get("draft_kv_cache_manager"), MiniMaxM3DraftSubpageView): + if isinstance(kwargs.get("draft_kv_cache_manager"), MiniMaxM3DraftKVCacheView): kwargs["draft_kv_cache_manager"] = None return super().add_dummy_requests(*args, **kwargs) @@ -639,11 +574,6 @@ def _kv_slot_geometry( kv_layout = self._main_kv_layout if kv_layout not in ("NHD", "HND"): raise ValueError(f"Unsupported kv_layout: {kv_layout}") - if self.is_fp8_subpaged_layer(layer_idx): - raise RuntimeError( - f"hybrid FP8 layer {layer_idx} uses four physical P32 pages; " - "use get_fp8_dense_buffers/get_dense_kv_subpage_pool instead" - ) if self.kv_cache_type == CacheTypeCpp.SELFKONLY: raise NotImplementedError( "MiniMaxM3KVCacheManagerV2 does not support the SELFKONLY cache type" @@ -719,23 +649,6 @@ def get_buffers( ``view[s, 0/1, ...]`` lands on this layer's K/V at slot ``s``. When omitted, ``kv_layout`` follows the selected sparse backend. """ - if self.is_fp8_subpaged_layer(layer_idx): - if kv_layout not in (None, "HND"): - raise ValueError( - "hybrid FP8 dense/Eagle buffers have a physical P32 HND layout; " - f"requested {kv_layout}" - ) - k, _v, slot_stride, pages_per_role = self._fp8_dense_data_buffers(layer_idx) - num_slots, _pages, num_heads, page_size, head_dim = k.shape - full = convert_to_torch_tensor( - TensorWrapper( - k.data_ptr(), - k.dtype, - [num_slots, slot_stride, num_heads, page_size, head_dim], - ) - ) - return full[:, : 2 * pages_per_role].unflatten(1, (2, pages_per_role)) - addr_key, torch_dtype, num_slots, scale, page_shape = self._kv_slot_geometry( layer_idx, kv_layout ) @@ -813,84 +726,6 @@ def get_block_scale_buffers( full_view = convert_to_torch_tensor(TensorWrapper(addr_key, torch_dtype, full_slot_shape)) return full_view[:, :2] - def _fp8_dense_data_buffers( - self, layer_idx: int - ) -> Tuple[torch.Tensor, torch.Tensor, int, int]: - """Return hybrid FP8 K/V views backed by physical P32 pages. - - Each logical P128 role page is laid out as four consecutive P32 HND - pages. ``slot_stride`` is measured in those physical pages and already - includes the V2 converter's expansion factor. - """ - if not self.is_fp8_subpaged_layer(layer_idx): - raise RuntimeError(f"layer {layer_idx} is not a physical-P32 hybrid FP8 layer") - local_layer_idx = self.layer_offsets[layer_idx] - physical_page = self.nvfp4_dense_tokens_per_block - pages_per_role = self.tokens_per_block // physical_page - addr_key = self.impl.get_mem_pool_base_address(local_layer_idx, Role.KEY) - addr_value = self.impl.get_mem_pool_base_address(local_layer_idx, Role.VALUE) - page_stride_key = self.impl.get_page_stride(local_layer_idx, Role.KEY) - page_stride_value = self.impl.get_page_stride(local_layer_idx, Role.VALUE) - assert page_stride_key == page_stride_value - assert addr_key + pages_per_role * page_stride_key == addr_value, ( - "M3 hybrid FP8 storage requires V immediately after K's physical " - f"P{physical_page} pages; layer={layer_idx} K={addr_key} " - f"stride={page_stride_key} V={addr_value}." - ) - - converter = self.impl.get_page_index_converter(local_layer_idx, Role.KEY) - assert int(converter.expansion) == pages_per_role, ( - f"layer {layer_idx} expected V2 expansion {pages_per_role}, got " - f"{int(converter.expansion)}" - ) - slot_stride = int(converter.scale) * pages_per_role - layer_offset_pages = int(converter.layer_offset) * pages_per_role - page_upper = self.impl.get_page_index_upper_bound(local_layer_idx, Role.KEY) - total_pages = int(page_upper) + layer_offset_pages - assert total_pages % slot_stride == 0 - num_slots = total_pages // slot_stride - - num_heads = self.num_kv_heads_per_layer[local_layer_idx] - head_dim = self.head_dim_per_layer[local_layer_idx] - full = convert_to_torch_tensor( - TensorWrapper( - addr_key, - torch.float8_e4m3fn, - [num_slots, slot_stride, num_heads, physical_page, head_dim], - ) - ) - return ( - full[:, :pages_per_role], - full[:, pages_per_role : 2 * pages_per_role], - slot_stride, - pages_per_role, - ) - - def get_fp8_dense_buffers(self, layer_idx: int) -> Tuple[torch.Tensor, torch.Tensor]: - """Return physical-P32 hybrid FP8 views as ``K, V``.""" - k, v, _slot_stride, _pages_per_role = self._fp8_dense_data_buffers(layer_idx) - return k, v - - def get_dense_kv_subpage_pool(self, layer_idx: int) -> Tuple[torch.Tensor, int, int]: - """Flat dense-attention pool, slot stride, and pages per K/V role.""" - if not self.is_fp8_subpaged_layer(layer_idx): - pool, slot_stride = self.get_kv_subpage_pool(layer_idx, "HND") - return pool, slot_stride, 1 - k, _v, slot_stride, pages_per_role = self._fp8_dense_data_buffers(layer_idx) - num_slots, _pages, num_heads, page_size, head_dim = k.shape - num_pages = (num_slots - 1) * slot_stride + 2 * pages_per_role - addr = k.data_ptr() - pool = convert_to_torch_tensor( - TensorWrapper(addr, k.dtype, [num_pages, num_heads, page_size, head_dim]) - ) - return pool, slot_stride, pages_per_role - - def get_dense_kv_scale_subpage_pool(self, layer_idx: int) -> Tuple[torch.Tensor, int, int]: - """Dense/Eagle layers are FP8 in the hybrid cache and have no scales.""" - raise RuntimeError( - f"hybrid FP8 dense/Eagle layer {layer_idx} has no NVFP4 block-scale pool" - ) - def get_kv_subpage_pool( self, layer_idx: int, kv_layout: str = "HND" ) -> Tuple[torch.Tensor, int]: @@ -908,11 +743,6 @@ def get_kv_subpage_pool( spanning ``num_slots * scale``, which would run off the pool by whatever this layer's K offset is inside a slot. """ - if self.is_fp8_subpaged_layer(layer_idx): - if kv_layout != "HND": - raise ValueError("hybrid FP8 dense/Eagle sub-pages are HND only") - pool, slot_stride, _pages_per_role = self.get_dense_kv_subpage_pool(layer_idx) - return pool, slot_stride addr_key, torch_dtype, num_slots, scale, page_shape = self._kv_slot_geometry( layer_idx, kv_layout ) @@ -928,8 +758,7 @@ def get_kv_scale_subpage_pool( """Return the flat NVFP4 scale pool paired with ``get_kv_subpage_pool``. The returned factor must match the packed-data factor so one K/V block - table addresses both pools. Token-size subdivision, when needed by the - Eagle draft view, is expressed by that view's expanded block table. + table addresses both pools. """ addr_key, torch_dtype, num_slots, scale, page_shape = self._kv_scale_slot_geometry( layer_idx, kv_layout @@ -1044,99 +873,72 @@ def get_block_ids_per_seq(self, request_ids): return padded_tensor -class MiniMaxM3DraftSubpageView: - """Present the shared manager's draft-layer pool at a smaller kernel page size. - - With unified KV cache the drafter's KV lives inside the shared manager's - 128-token logical blocks, but the Eagle3 kernels are only healthy at - 32-token pages on this architecture. This view flows wherever a separate - draft manager would (``get_draft_kv_cache_manager`` and the - attention-metadata draft swap) and re-expresses the geometry only: - - * the single pool pointer is re-rooted at the drafter's K address, and - the draft layer's row in the pool mapping points at that pool; - * the block table expands each logical slot ``s`` into sub-pages — - K at ``s*scale*subdiv + j`` (``j < subdiv``), V at ``+subdiv`` - (``scale`` = the drafter's pages per mega-slot, from - ``_kv_slot_geometry``; same layout trick as the dense-layer - trtllm-gen adapter). - - The attention op reads ``tokens_per_block`` and the pool pointers from - the metadata's manager, so no attention-backend changes are needed. The - view owns no blocks: lifecycle stays entirely with the shared manager. +class MiniMaxM3DraftKVCacheView: + """Present the shared manager's dense draft K/V pages to TRTLLM-Gen. - Retires with the P128 Eagle kernel fixes; see the retirement plan on - ``MiniMaxM3KVCacheManagerV2.draft_manager_tokens_per_block``. + MiniMax-M3 stores every layer in one non-uniform mega-slot: sparse target + layers add an index-K page, while the dense draft layer contributes only K + and V. This view roots the pool at the draft K page and maps logical slot + ``s`` to ``K=s*scale, V=K+1``. It owns no block lifecycle. """ - def __init__(self, manager, draft_layer_ids: Sequence[int], subpage_tokens: int): + # P128 is exported for this view's exact dense-GQA shapes, but not for all + # shapes in the TRTLLM-Gen artifact. + trtllm_gen_extra_tokens_per_block = frozenset({128}) + + def __init__(self, manager, draft_layer_ids: Sequence[int]): self._manager = manager - self.tokens_per_block = int(subpage_tokens) - layer_id = draft_layer_ids[0] - is_hybrid_fp8 = bool( - getattr(manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_id) - ) - if is_hybrid_fp8: - assert self.tokens_per_block == manager.nvfp4_dense_tokens_per_block, ( - "hybrid FP8 Eagle draft attention must use the physical dense-cache " - f"page size P{manager.nvfp4_dense_tokens_per_block}, got " - f"P{self.tokens_per_block}" + if len(draft_layer_ids) != 1: + raise ValueError( + "MiniMax-M3's native P128 draft view supports exactly one " + f"draft layer, got {len(draft_layer_ids)}" ) - assert manager.tokens_per_block % self.tokens_per_block == 0, ( - f"subpage size {subpage_tokens} must divide manager " - f"tokens_per_block {manager.tokens_per_block}" - ) - self._subdiv = manager.tokens_per_block // self.tokens_per_block - # The hybrid layout puts dense/Eagle FP8 pages in a separate physical - # pool from sparse NVFP4 data/scales. Root this single-pool view at the - # draft layer's K address, while sourcing raw logical slot IDs from the - # draft layer's actual V2 pool. - local = manager.layer_offsets[layer_id] - self._source_pool_id = int(manager.kv_cache_pool_mapping[int(local), 0]) - if is_hybrid_fp8: - k, _v, slot_stride, pages_per_role = manager._fp8_dense_data_buffers(layer_id) - assert pages_per_role == self._subdiv - addr_key = k.data_ptr() - self._num_slots = int(k.shape[0]) - self._slot_units = slot_stride + if manager.tokens_per_block != 128: + raise ValueError( + "MiniMax-M3's native P128 draft view requires tokens_per_block=128, " + f"got {manager.tokens_per_block}" + ) + if manager.enable_swa_scratch_reuse: + raise ValueError( + "MiniMax-M3's native P128 draft view does not support SWA scratch reuse" + ) + + layer_id = draft_layer_ids[0] + local_layer_id = manager.layer_offsets[layer_id] + source_pool_id = int(manager.kv_cache_pool_mapping[int(local_layer_id), 0]) + + flat_pool, slot_stride = manager.get_kv_subpage_pool(layer_id, "HND") + if int(manager.kv_offset[source_pool_id]) != 1 or manager._stream is None: + raise ValueError("MiniMax-M3's native P128 draft block-table mapping is unavailable") + + self._source_pool_id = source_pool_id + self._flat_pool = flat_pool + if manager.is_fp8_dense_layer(layer_id): self.dtype = DataType.FP8 - else: - addr_key, _dt, num_slots, scale, _shape = manager._kv_slot_geometry(layer_id, None) - self._num_slots = int(num_slots) - self._slot_units = scale * self._subdiv self.num_pools = 1 self.num_attention_op_pools = 1 - self.max_blocks_per_seq = manager.max_blocks_per_seq * self._subdiv - # Single-pool pointer rooted at the drafter's K; the op derives the - # page stride from tokens_per_block, so unit indices below address - # 32-token drafter pages directly. + self.max_blocks_per_seq = manager.max_blocks_per_seq self.kv_cache_pool_pointers = torch.tensor( - [[addr_key, 0]], dtype=torch.int64, pin_memory=prefer_pinned() - ) - mapping = manager.kv_cache_pool_mapping.clone() - mapping[int(local)] = torch.tensor([0, 0], dtype=mapping.dtype) - self.kv_cache_pool_mapping = mapping - # Placeholder host mirror: the dense TRTLLM path plans from device - # offsets; nothing reads the host table during the draft window. - self.host_kv_cache_block_offsets = torch.zeros( - (1, 1, 2, 1), dtype=torch.int32, pin_memory=prefer_pinned() - ) - self._slots_host: Optional[np.ndarray] = None - self._arange: Optional[torch.Tensor] = None + [[flat_pool.data_ptr(), 0]], dtype=torch.int64, pin_memory=prefer_pinned() + ) + self.kv_cache_pool_mapping = manager.kv_cache_pool_mapping.clone() + self.kv_cache_pool_mapping[int(local_layer_id)] = 0 + # The view is rooted at this layer's K page, so its block-table scale + # is the layer-local stride of ``flat_pool``. The source pool's scale + # is measured in the first layer's page units and can differ for a + # heterogeneous NVFP4 mega-slot (for example, 171 versus 8 here). + self.index_scales = torch.tensor( + [slot_stride], dtype=torch.int32, pin_memory=prefer_pinned() + ) + self.kv_offset = manager.kv_offset[source_pool_id : source_pool_id + 1] + self.host_kv_cache_block_offsets = manager.host_kv_cache_block_offsets[ + source_pool_id : source_pool_id + 1 + ] @property def blocks_in_primary_pool(self) -> int: - """Flattened sub-page index bound relative to the draft K pointer. - - ``FlashInferTrtllmGenFmha`` uses this value to size the flat paged-KV - tensor passed to FlashInfer. The wrapped V2 manager reports its bound - in 128-token page units relative to a different pool base, so - delegating that property through ``__getattr__`` under-describes this - 32-token, draft-K-rooted view. The final slot contributes only this - layer's K and V pages; inter-layer padding after V is not addressable - from the view and need not be included. - """ - return (self._num_slots - 1) * self._slot_units + 2 * self._subdiv + """Return the flat P128 page bound relative to the draft K pointer.""" + return int(self._flat_pool.shape[0]) def __getattr__(self, name): manager = self.__dict__.get("_manager") @@ -1151,56 +953,6 @@ def host_kv_cache_pool_pointers(self): def free_resources(self, request) -> None: """No-op: block lifecycle belongs to the shared manager.""" - def _host_block_table( - self, - slot_rows: Sequence[Sequence[int]], - num_seqs: int, - max_slots: int, - dtype: torch.dtype, - ) -> torch.Tensor: - """Expand this batch's slot ids into a freshly allocated pinned table. - - The buffer is allocated per call on purpose: it is the source of an - asynchronous H2D copy, whose source is read at copy execution time - rather than enqueue time. A persistent buffer refilled in place would - let the next iteration's refill clobber a still-pending copy — the - drafter would then index another batch's blocks (nvbug 6293536, whose - rationale the V1 manager spells out on - ``KVCacheManager._stage_block_offsets_for_copy``). The caching host - allocator keeps this block alive until the copy retires. The numpy - scratch below stays persistent: it is only ever read synchronously. - """ - sub = self._subdiv - if ( - self._slots_host is None - or self._slots_host.shape[0] < num_seqs - or self._slots_host.shape[1] != max_slots - ): - self._slots_host = np.zeros((num_seqs, max_slots), dtype=np.int32) - self._arange = torch.arange(sub, dtype=dtype) - slots_np = self._slots_host[:num_seqs] - slots_np.fill(0) - # Ragged fill is the only per-row work (numpy parses each row's list at - # C speed); the arithmetic below is one fused expansion across the - # batch. Pad/BAD_PAGE_INDEX entries clamp to slot 0 (safe pages: - # kernels never read past kv_lens). - for i, row in enumerate(slot_rows[:num_seqs]): - n = min(len(row), max_slots) - if n > 0: - slots_np[i, :n] = row[:n] - np.clip(slots_np, 0, None, out=slots_np) - slots = torch.from_numpy(slots_np).to(dtype) - host = torch.empty( - (num_seqs, 2, max_slots * sub), - dtype=dtype, - pin_memory=prefer_pinned(), - device="cpu", - ) - out = host.view(num_seqs, 2, max_slots, sub) - torch.add(slots.unsqueeze(-1) * self._slot_units, self._arange, out=out[:, 0]) - torch.add(out[:, 0], sub, out=out[:, 1]) - return host - def copy_batch_block_offsets( self, dst_tensor: torch.Tensor, @@ -1210,14 +962,19 @@ def copy_batch_block_offsets( num_seqs: int, max_blocks: Optional[int] = None, ) -> None: - # Raw logical slot ids from the draft layer's physical pool. - slot_rows = self._manager._get_batch_cache_indices_by_pool_id( - request_ids, pool_id=self._source_pool_id - ) - host = self._host_block_table( - slot_rows, num_seqs, dst_tensor.shape[-1] // self._subdiv, dst_tensor.dtype + # Call the V2 implementation unbound so it reads this view's + # ``index_scales``, ``kv_offset`` and ``host_kv_cache_block_offsets`` + # rather than the shared manager's; ``__getattr__`` would bind the + # manager's own copy otherwise. + KVCacheManagerV2.copy_batch_block_offsets( + self, + dst_tensor, + request_ids, + beam_width, + num_contexts, + num_seqs, + max_blocks=max_blocks, ) - dst_tensor[0, :num_seqs].copy_(host, non_blocking=True) def get_minimax_m3_kv_cache_manager_cls(): diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_backend.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_backend.py index 6dbace99fed1..9ada707cfe75 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_backend.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_backend.py @@ -53,7 +53,7 @@ ) from .trtllm_gen_dense_decode import ( dense_decode_unsupported_reason, - uniform_dense_subpage_geometry, + uniform_dense_subpages_per_slot, write_subpage_block_table, ) @@ -312,7 +312,6 @@ class MiniMaxM3MsaSparseAttentionMetadata(TrtllmAttentionMetadata): # factor, or 0 where the pool has no single one; see msa_subpage_rows. msa_subpage_block_table: Optional[torch.Tensor] = None _msa_subpages_per_slot: int = 0 - _msa_pages_per_dense_role: int = 1 # Per-request kv_lens as staged by prepare(), before the overlap scheduler # corrects them. on_update_kv_lens clamps against this; see there. msa_kv_lens_staged: Optional[torch.Tensor] = None @@ -598,16 +597,14 @@ def _create_msa_buffers(self) -> None: ) # Resolved once here rather than per step: the factor is fixed by the # pool's layout for the life of the manager. - self._msa_subpages_per_slot, self._msa_pages_per_dense_role = ( - uniform_dense_subpage_geometry(kv_cache_manager) - ) + self._msa_subpages_per_slot = uniform_dense_subpages_per_slot(kv_cache_manager) if self._msa_subpages_per_slot > 0: self.msa_subpage_block_table = self.get_empty( buffers, ( max_num_sequences, 2, - max_blocks_per_seq * self._msa_pages_per_dense_role, + max_blocks_per_seq, ), cache_name="msa_subpage_block_table", dtype=torch.int32, @@ -1459,7 +1456,6 @@ def _build_msa_fields(self) -> None: self.msa_block_table[:batch_size], self._msa_subpages_per_slot, self.msa_subpage_block_table[:batch_size], - self._msa_pages_per_dense_role, ) # Staging for on_update_kv_lens. @@ -1530,11 +1526,7 @@ def msa_write_layer_caches( Requires prepared metadata (msa_out_cache_loc filled), the same contract as the writes it replaces. """ - from .msa_scatter import ( - fused_write_layer_caches, - fused_write_layer_caches_nvfp4, - fused_write_subpaged_layer_caches, - ) + from .msa_scatter import fused_write_layer_caches, fused_write_layer_caches_nvfp4 idx_cache = self.msa_idx_k_cache(layer_idx) if idx_k is not None else None num_tokens = int(k.shape[0]) @@ -1542,9 +1534,6 @@ def msa_write_layer_caches( is_nvfp4_layer = getattr(self.kv_cache_manager, "is_nvfp4_layer", lambda _layer_idx: False)( layer_idx ) - is_fp8_subpaged_layer = getattr( - self.kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False - )(layer_idx) if is_nvfp4_layer: if kv_scale_orig_quant is None: raise RuntimeError( @@ -1572,13 +1561,6 @@ def msa_write_layer_caches( "MiniMax-M3 NVFP4 cache writer requires CUDA HND P128/P32 cache views, " "contiguous logical K/V rows, and FP32 K/V quantization scales" ) - elif is_fp8_subpaged_layer: - k_view, v_view = self.kv_cache_manager.get_fp8_dense_buffers(layer_idx) - if not fused_write_subpaged_layer_caches(k_view, v_view, out_cache_loc, k, v): - raise RuntimeError( - "MiniMax-M3 hybrid FP8 dense/Eagle cache writer requires CUDA " - "HND P32 sub-page views and contiguous logical K/V rows" - ) else: buffers = self.kv_cache_manager.get_buffers(layer_idx, kv_layout="HND") k_view, v_view = buffers[:, 0], buffers[:, 1] diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_scatter.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_scatter.py index 2693d479cb76..899939eed973 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_scatter.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_scatter.py @@ -71,61 +71,6 @@ def _fused_paged_scatter_kernel( tl.store(i_dst, i_vals.to(idx_cache.dtype.element_ty), mask=valid) -@triton.jit -def _fused_subpaged_scatter_kernel( - k_src, - v_src, - k_cache, - v_cache, - out_cache_loc, - k_src_row_stride, - v_src_row_stride, - kc_stride_page, - kc_stride_subpage, - kc_stride_head, - kc_stride_tok, - vc_stride_page, - vc_stride_subpage, - vc_stride_head, - vc_stride_tok, - logical_tokens_per_block, - physical_tokens_per_block, - H: tl.constexpr, - D: tl.constexpr, -): - """Scatter FP8 K/V into P32 pages inside one logical P128 slot.""" - t = tl.program_id(0).to(tl.int64) - slot = tl.load(out_cache_loc + t).to(tl.int64) - valid = slot >= 0 - page = slot // logical_tokens_per_block - logical_within = slot % logical_tokens_per_block - subpage = logical_within // physical_tokens_per_block - within = logical_within % physical_tokens_per_block - d = tl.arange(0, D) - for h in tl.static_range(H): - src = t * k_src_row_stride + h * D + d - k_vals = tl.load(k_src + src) - v_vals = tl.load(v_src + t * v_src_row_stride + h * D + d) - k_dst = ( - k_cache - + page * kc_stride_page - + subpage * kc_stride_subpage - + h * kc_stride_head - + within * kc_stride_tok - + d - ) - v_dst = ( - v_cache - + page * vc_stride_page - + subpage * vc_stride_subpage - + h * vc_stride_head - + within * vc_stride_tok - + d - ) - tl.store(k_dst, k_vals.to(k_cache.dtype.element_ty), mask=valid) - tl.store(v_dst, v_vals.to(v_cache.dtype.element_ty), mask=valid) - - @triton.jit def _fused_nvfp4_paged_scatter_kernel( k_data_src, @@ -320,62 +265,6 @@ def fused_write_layer_caches( return True -def fused_write_subpaged_layer_caches( - k_cache: torch.Tensor, - v_cache: torch.Tensor, - out_cache_loc: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, -) -> bool: - """Write ordinary K/V into physical sub-pages of a logical cache block. - - Hybrid M3 stores dense and shared-Eagle K/V as FP8 P32 pages while its - allocator and request lifecycle remain P128. ``k_cache``/``v_cache`` are - ``[logical_slots, pages_per_role, heads, P32, D]`` zero-copy views. - """ - if not (k.is_cuda and k_cache.is_cuda): - return False - if k_cache.dim() != 5 or v_cache.shape != k_cache.shape: - return False - if k_cache.stride(-1) != 1 or v_cache.stride(-1) != 1: - return False - _num_slots, pages_per_role, num_heads, physical_page, head_dim = k_cache.shape - if pages_per_role <= 0 or physical_page <= 0 or (head_dim & (head_dim - 1)) != 0: - return False - inner = num_heads * head_dim - k_stride = _row_stride_if_fusable(k, inner) - v_stride = _row_stride_if_fusable(v, inner) - if k_stride is None or v_stride is None: - return False - num_tokens = int(out_cache_loc.shape[0]) - if num_tokens == 0: - return True - - _fused_subpaged_scatter_kernel[(num_tokens,)]( - k, - v, - k_cache, - v_cache, - out_cache_loc, - k_stride, - v_stride, - k_cache.stride(0), - k_cache.stride(1), - k_cache.stride(2), - k_cache.stride(3), - v_cache.stride(0), - v_cache.stride(1), - v_cache.stride(2), - v_cache.stride(3), - pages_per_role * physical_page, - physical_page, - H=num_heads, - D=head_dim, - num_warps=2, - ) - return True - - def fused_write_layer_caches_nvfp4( k_data_cache: torch.Tensor, v_data_cache: torch.Tensor, @@ -527,5 +416,4 @@ def fused_write_layer_caches_nvfp4( __all__ = [ "fused_write_layer_caches", "fused_write_layer_caches_nvfp4", - "fused_write_subpaged_layer_caches", ] diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_utils.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_utils.py index 47502b7faa59..b2d3b91164e1 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_utils.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/msa_utils.py @@ -203,11 +203,6 @@ def msa_paged_kv(kv_cache_manager, layer_idx: int) -> Tuple[torch.Tensor, torch. runtime and needs only each page's [page_size, head_dim] block to be contiguous, which this view satisfies, so no copy is required. """ - if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx): - raise RuntimeError( - "hybrid FP8 dense/Eagle cache is physically P32; use the direct " - "TRTLLM-Gen dense adapter instead of msa_paged_kv" - ) buffers = kv_cache_manager.get_buffers(layer_idx, kv_layout="HND") return buffers[:, 0], buffers[:, 1] @@ -225,16 +220,6 @@ def write_msa_main_kv( resident before the sparse GQA runs. The write uses the head-major HND view so `msa_paged_kv` can return a zero-copy view. """ - if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx): - from .msa_scatter import fused_write_subpaged_layer_caches - - k_view, v_view = kv_cache_manager.get_fp8_dense_buffers(layer_idx) - if not fused_write_subpaged_layer_caches(k_view, v_view, out_cache_loc, k, v): - raise RuntimeError( - "MiniMax-M3 hybrid FP8 dense/Eagle cache write requires CUDA " - "P32 sub-page views and contiguous K/V rows" - ) - return buffers = kv_cache_manager.get_buffers(layer_idx, kv_layout="HND") k_view, v_view = buffers[:, 0], buffers[:, 1] num_kv_heads = int(k_view.shape[1]) diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/trtllm_gen_dense_decode.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/trtllm_gen_dense_decode.py index 6400259ee5dc..5ec0906b8a34 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/trtllm_gen_dense_decode.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/trtllm_gen_dense_decode.py @@ -6,8 +6,7 @@ they only run there because MsaSparseGqaFmha claims every M3 target layer. MSA's kernel uses the context schedule, spending a 128-row Q tile on one decode token, while trtllm-gen has generation scheduling for exactly this -shape. The NVFP4 path also uses trtllm-gen context attention because its -shipped cubins consume physical P32 packed data and block-scale pages. +shape. FlashInferTrtllmGenFmha cannot be reused as-is. It reaches the pool through build_trtllm_gen_kv_cache_metadata, which assumes each layer contributes @@ -28,7 +27,6 @@ import torch from tensorrt_llm._torch.memory_buffer_utils import get_memory_buffers -from tensorrt_llm.bindings import DataType from .msa_utils import check_decode_span_shape @@ -36,25 +34,11 @@ class _MiniMaxM3DenseKVCacheManager(Protocol): """Cache-pool surface consumed by the dense TRTLLM-gen helpers.""" - dtype: DataType layer_offsets: Mapping[int, int] sparse_layer_ids: Collection[int] - def is_nvfp4_layer(self, layer_idx: int) -> bool: ... - def get_kv_subpage_pool(self, layer_idx: int, kv_layout: str) -> tuple[torch.Tensor, int]: ... - def get_dense_kv_subpage_pool(self, layer_idx: int) -> tuple[torch.Tensor, int, int]: ... - - def get_dense_kv_scale_subpage_pool(self, layer_idx: int) -> tuple[torch.Tensor, int, int]: ... - - -def _layer_uses_nvfp4(kv_cache_manager: _MiniMaxM3DenseKVCacheManager, layer_idx: int) -> bool: - predicate = getattr(kv_cache_manager, "is_nvfp4_layer", None) - if predicate is not None: - return bool(predicate(layer_idx)) - return getattr(kv_cache_manager, "dtype", None) == DataType.NVFP4 - @functools.lru_cache(maxsize=None) def _counter_size(num_heads: int, max_num_requests: int, device_index: int) -> int: @@ -111,104 +95,27 @@ def _workspace(q_dtype: torch.dtype, num_heads: int, head_dim: int, num_kv_heads return int(layout["trtllm_gen_workspace_size"]) -@functools.lru_cache(maxsize=None) -def _context_workspace( - q_dtype: torch.dtype, - max_num_requests: int, - max_num_tokens: int, - num_heads: int, - head_dim: int, -) -> int: - from tensorrt_llm._torch.attention_backend.fmha.flashinfer_trtllm_gen import ( - _get_context_workspace_size, - ) - - return int( - _get_context_workspace_size( - q_dtype, - max_num_requests, - max_num_tokens, - num_heads, - head_dim, - 0, - True, - ) - ) - - def _dense_kv_inputs( q: torch.Tensor, kv_cache_manager: _MiniMaxM3DenseKVCacheManager, layer_idx: int, - *, - sm_scale: float, - kv_scale_quant_orig: Optional[torch.Tensor], -) -> tuple[ - torch.Tensor, - torch.Tensor, - Optional[torch.Tensor], - int, - int, - float | torch.Tensor, - float | torch.Tensor, -]: - """Resolve direct TRTLLM-gen inputs for one M3 dense layer.""" - get_dense_pool = getattr(kv_cache_manager, "get_dense_kv_subpage_pool", None) - if get_dense_pool is None: - kv_pool, subpages_per_slot = kv_cache_manager.get_kv_subpage_pool(layer_idx, "HND") - pages_per_role = 1 - else: - kv_pool, subpages_per_slot, pages_per_role = get_dense_pool(layer_idx) - kv_scale_pool = None - bmm1_scale: float | torch.Tensor = sm_scale - bmm2_scale: float | torch.Tensor = 1.0 - - if _layer_uses_nvfp4(kv_cache_manager, layer_idx): - if kv_scale_quant_orig is None: - raise RuntimeError("MiniMax-M3 dense NVFP4 attention requires [Q, K, V] scales") - if kv_scale_quant_orig.dtype != torch.float32 or kv_scale_quant_orig.numel() < 3: - raise ValueError("MiniMax-M3 dense NVFP4 scales must be FP32 [Q, K, V]") - kv_scale_pool, scale_factor, scale_pages = kv_cache_manager.get_dense_kv_scale_subpage_pool( - layer_idx - ) - if (int(scale_factor), int(scale_pages)) != ( - int(subpages_per_slot), - int(pages_per_role), - ): - raise RuntimeError( - "MiniMax-M3 NVFP4 data/scale pools have different block-table geometry" - ) +) -> tuple[torch.Tensor, torch.Tensor, int]: + """Resolve the query, flat K/V pool and slot stride for one M3 dense layer.""" + kv_pool, subpages_per_slot = kv_cache_manager.get_kv_subpage_pool(layer_idx, "HND") + if kv_pool.dtype == torch.float8_e4m3fn and q.dtype != torch.float8_e4m3fn: q = q.to(torch.float8_e4m3fn) - kv_pool = kv_pool.view(torch.uint8) - kv_scale_pool = kv_scale_pool.view(torch.float8_e4m3fn) - raw_bmm1 = kv_scale_quant_orig[1:2] * float(sm_scale) - bmm1_scale = torch.cat((raw_bmm1, raw_bmm1 * 1.4426950408889634)) - bmm2_scale = kv_scale_quant_orig[2:3] - elif kv_pool.dtype == torch.float8_e4m3fn and q.dtype != torch.float8_e4m3fn: - q = q.to(torch.float8_e4m3fn) - - return ( - q, - kv_pool, - kv_scale_pool, - int(subpages_per_slot), - int(pages_per_role), - bmm1_scale, - bmm2_scale, - ) + return q, kv_pool, int(subpages_per_slot) def subpage_block_table( block_table: torch.Tensor, subpages_per_slot: int, reserve: bool = False, - pages_per_role: int = 1, ) -> torch.Tensor: """Expand a slot table into trtllm-gen's separate K and V page rows. uses_shared_paged_kv_idx is False for TensorRT-LLM, so the kernel takes - [batch, 2, max_blocks] and indexes K and V independently. With physical - P32 NVFP4 pages, each logical slot expands to ``pages_per_role`` entries. + [batch, 2, max_blocks] and indexes K and V independently. The result is a function of the slot table and that factor alone, so every dense layer of a step would compute the same one. prepare() therefore @@ -218,12 +125,12 @@ def subpage_block_table( """ batch, max_blocks = block_table.shape out = get_memory_buffers().get_buffer( - [batch, 2, max_blocks * pages_per_role], + [batch, 2, max_blocks], torch.int32, buffer_name="m3_trtllm_gen_subpage_block_table", reserve_buffer=reserve, ) - write_subpage_block_table(block_table, subpages_per_slot, out, pages_per_role) + write_subpage_block_table(block_table, subpages_per_slot, out) return out @@ -231,17 +138,10 @@ def write_subpage_block_table( block_table: torch.Tensor, subpages_per_slot: int, out: torch.Tensor, - pages_per_role: int = 1, ) -> None: """Write the K and V sub-page rows of block_table into out.""" - if pages_per_role == 1: - torch.mul(block_table, subpages_per_slot, out=out[:, 0]) - torch.add(out[:, 0], 1, out=out[:, 1]) - return - offsets = torch.arange(pages_per_role, dtype=block_table.dtype, device=block_table.device) - key_pages = (block_table.unsqueeze(-1) * subpages_per_slot + offsets).flatten(1) - out[:, 0].copy_(key_pages) - torch.add(out[:, 0], pages_per_role, out=out[:, 1]) + torch.mul(block_table, subpages_per_slot, out=out[:, 0]) + torch.add(out[:, 0], 1, out=out[:, 1]) def uniform_subpages_per_slot( @@ -262,14 +162,14 @@ def uniform_subpages_per_slot( return factors.pop() if len(factors) == 1 else 0 -def uniform_dense_subpage_geometry( +def uniform_dense_subpages_per_slot( kv_cache_manager: _MiniMaxM3DenseKVCacheManager, -) -> tuple[int, int]: - """Common dense slot stride and pages per role, or ``(0, 0)``.""" - get_pool = getattr(kv_cache_manager, "get_dense_kv_subpage_pool", None) +) -> int: + """Common dense slot stride, or 0 when dense layer groups disagree.""" + get_pool = getattr(kv_cache_manager, "get_kv_subpage_pool", None) layer_offsets = getattr(kv_cache_manager, "layer_offsets", None) - if not layer_offsets: - return 0, 0 + if get_pool is None or not layer_offsets: + return 0 dense_layers = [ layer_idx for layer_idx in layer_offsets @@ -277,18 +177,9 @@ def uniform_dense_subpage_geometry( or layer_idx not in kv_cache_manager.sparse_layer_ids ] if not dense_layers: - return 0, 0 - if get_pool is None: - legacy = getattr(kv_cache_manager, "get_kv_subpage_pool", None) - if legacy is None: - return 0, 0 - geometries = {(int(legacy(layer_idx, "HND")[1]), 1) for layer_idx in dense_layers} - else: - geometries = set() - for layer_idx in dense_layers: - _pool, slot_stride, pages_per_role = get_pool(layer_idx) - geometries.add((int(slot_stride), int(pages_per_role))) - return geometries.pop() if len(geometries) == 1 else (0, 0) + return 0 + strides = {int(get_pool(layer_idx, "HND")[1]) for layer_idx in dense_layers} + return strides.pop() if len(strides) == 1 else 0 def minimax_m3_trtllm_gen_dense_decode( @@ -305,7 +196,6 @@ def minimax_m3_trtllm_gen_dense_decode( max_num_requests: int, staged_subpage_table: Optional[torch.Tensor] = None, staged_subpages_per_slot: int = 0, - kv_scale_quant_orig: Optional[torch.Tensor] = None, enable_pdl: bool = True, ) -> None: """Full-context decode attention through trtllm-gen, in place into output. @@ -327,21 +217,7 @@ def minimax_m3_trtllm_gen_dense_decode( decode_query_len, ) - ( - q, - kv_pool, - kv_scale_pool, - subpages_per_slot, - pages_per_role, - bmm1_scale, - bmm2_scale, - ) = _dense_kv_inputs( - q, - kv_cache_manager, - layer_idx, - sm_scale=sm_scale, - kv_scale_quant_orig=kv_scale_quant_orig, - ) + q, kv_pool, subpages_per_slot = _dense_kv_inputs(q, kv_cache_manager, layer_idx) num_heads = int(q.shape[1]) reserve = torch.cuda.is_current_stream_capturing() @@ -352,9 +228,7 @@ def minimax_m3_trtllm_gen_dense_decode( reserve_buffer=reserve, ) if staged_subpage_table is None or staged_subpages_per_slot != subpages_per_slot: - staged_subpage_table = subpage_block_table( - block_table, subpages_per_slot, reserve, pages_per_role - ) + staged_subpage_table = subpage_block_table(block_table, subpages_per_slot, reserve) _trtllm_gen_batch_decode_with_kv_cache( q, # query @@ -366,8 +240,8 @@ def minimax_m3_trtllm_gen_dense_decode( staged_subpage_table, # block_tables seq_lens, # seq_lens max_seq_len, # max_seq_len - bmm1_scale, # bmm1_scale - bmm2_scale, # bmm2_scale + sm_scale, # bmm1_scale + 1.0, # bmm2_scale -1, # window_left: M3 dense layers are fully causal output, # out None, # sinks @@ -375,160 +249,11 @@ def minimax_m3_trtllm_gen_dense_decode( decode_query_len, # q_len_per_req None, # max_q_len None, # cum_seq_lens_q - kv_scale_pool, # NVFP4 E4M3 block scales, otherwise None + None, # kv_scale_pool: dense layers use FP8/P128 False, # uses_shared_paged_kv_idx ) -def minimax_m3_trtllm_gen_dense_context( - q: torch.Tensor, - kv_cache_manager: _MiniMaxM3DenseKVCacheManager, - layer_idx: int, - block_table: torch.Tensor, - seq_lens: torch.Tensor, - cu_q_lens: torch.Tensor, - cu_kv_lens: torch.Tensor, - *, - sm_scale: float, - output: torch.Tensor, - max_q_len: int, - max_kv_len: int, - max_num_requests: int, - kv_scale_quant_orig: Optional[torch.Tensor] = None, - staged_subpage_table: Optional[torch.Tensor] = None, - staged_subpages_per_slot: int = 0, - enable_pdl: bool = True, -) -> None: - """Full-context attention for M3 dense layers, including NVFP4 KV.""" - from tensorrt_llm._torch.attention_backend.fmha.flashinfer_trtllm_gen import ( - _trtllm_gen_batch_context_with_kv_cache, - ) - - ( - q, - kv_pool, - kv_scale_pool, - subpages_per_slot, - pages_per_role, - bmm1_scale, - bmm2_scale, - ) = _dense_kv_inputs( - q, - kv_cache_manager, - layer_idx, - sm_scale=sm_scale, - kv_scale_quant_orig=kv_scale_quant_orig, - ) - batch_size = int(seq_lens.shape[0]) - num_heads = int(q.shape[1]) - reserve = torch.cuda.is_current_stream_capturing() - workspace = get_memory_buffers().get_buffer( - [ - _context_workspace( - q.dtype, max_num_requests, int(q.shape[0]), num_heads, int(q.shape[2]) - ) - ], - torch.uint8, - buffer_name="m3_trtllm_gen_context_workspace", - reserve_buffer=reserve, - ) - if staged_subpage_table is None or staged_subpages_per_slot != subpages_per_slot: - staged_subpage_table = subpage_block_table( - block_table, subpages_per_slot, reserve, pages_per_role - ) - - _trtllm_gen_batch_context_with_kv_cache( - q, - kv_pool, - workspace, - _counter_buffer(q.device, num_heads, max_num_requests, reserve), - staged_subpage_table, - seq_lens, - max_q_len, - max_kv_len, - bmm1_scale, - bmm2_scale, - batch_size, - cu_q_lens, - cu_kv_lens, - -1, - output, - None, - enable_pdl, - kv_scale_pool, - False, - True, - ) - - -def minimax_m3_trtllm_gen_dense_attention( - q: torch.Tensor, - kv_cache_manager: _MiniMaxM3DenseKVCacheManager, - layer_idx: int, - metadata, - *, - sm_scale: float, - output: torch.Tensor, - kv_scale_quant_orig: Optional[torch.Tensor] = None, -) -> None: - """Dispatch the packed M3 dense batch by context/generation phase.""" - qo_lens = metadata.msa_qo_lens_cpu - kv_lens = metadata.msa_kv_lens_cpu - if qo_lens is None or kv_lens is None: - raise RuntimeError("MiniMax-M3 dense attention metadata was not prepared") - num_contexts = int(metadata.num_contexts or 0) - batch_size = int(qo_lens.shape[0]) - ctx_tokens = int(qo_lens[:num_contexts].sum().item()) if num_contexts else 0 - - staged_rows = getattr(metadata, "msa_subpage_rows", None) - if num_contexts: - staged_table, staged_factor = ( - staged_rows(0, num_contexts) if staged_rows is not None else (None, 0) - ) - minimax_m3_trtllm_gen_dense_context( - q[:ctx_tokens], - kv_cache_manager, - layer_idx, - metadata.msa_block_table[:num_contexts], - metadata.msa_seq_lens_cuda[:num_contexts], - metadata.msa_cu_q_lens[: num_contexts + 1], - metadata.msa_cu_kv_lens[: num_contexts + 1], - sm_scale=sm_scale, - output=output[:ctx_tokens], - max_q_len=int(qo_lens[:num_contexts].max().item()), - max_kv_len=int(kv_lens[:num_contexts].max().item()), - max_num_requests=int(metadata.max_num_requests), - kv_scale_quant_orig=kv_scale_quant_orig, - staged_subpage_table=staged_table, - staged_subpages_per_slot=staged_factor, - ) - - if num_contexts == batch_size: - return - gen_qo = qo_lens[num_contexts:] - decode_query_len = int(gen_qo[0].item()) - if not torch.equal(gen_qo, torch.full_like(gen_qo, decode_query_len)): - raise NotImplementedError("MiniMax-M3 dense generation requires a uniform query length") - staged_table, staged_factor = ( - staged_rows(num_contexts, batch_size) if staged_rows is not None else (None, 0) - ) - minimax_m3_trtllm_gen_dense_decode( - q[ctx_tokens:], - kv_cache_manager, - layer_idx, - metadata.msa_block_table[num_contexts:batch_size], - metadata.msa_seq_lens_cuda[num_contexts:batch_size], - sm_scale=sm_scale, - output=output[ctx_tokens:], - decode_query_len=decode_query_len, - max_seq_len=int(kv_lens[num_contexts:].max().item()), - max_num_requests=int(metadata.max_num_requests), - staged_subpage_table=staged_table, - staged_subpages_per_slot=staged_factor, - kv_scale_quant_orig=kv_scale_quant_orig, - ) - - @functools.lru_cache(maxsize=1) def _flashinfer_available() -> bool: """Whether flashinfer can be imported, resolved once for the process. @@ -555,23 +280,6 @@ def dense_decode_unsupported_reason( return "the KV cache manager does not expose a flat sub-page pool." if int(head_dim) != 128: return f"head_dim {int(head_dim)}; only 128 has trtllm-gen H128 cubins." - dense_layers = [ - layer_idx - for layer_idx in getattr(kv_cache_manager, "layer_offsets", {}) - if layer_idx not in getattr(kv_cache_manager, "sparse_layer_ids", ()) - ] - probe_layer = dense_layers[0] if dense_layers else 0 - if _layer_uses_nvfp4(kv_cache_manager, probe_layer): - if not hasattr(kv_cache_manager, "get_dense_kv_scale_subpage_pool"): - return "the NVFP4 manager does not expose a dense block-scale pool." - if not dense_layers: - return "the NVFP4 manager has no dense attention layer." - dense_pool, _stride, _pages = kv_cache_manager.get_dense_kv_subpage_pool(dense_layers[0]) - if int(dense_pool.shape[2]) != 32: - return ( - f"the NVFP4 dense pool uses P{int(dense_pool.shape[2])}; " - "the shipped trtllm-gen cubins require P32." - ) if not _flashinfer_available(): return "flashinfer is not installed." return None @@ -579,11 +287,9 @@ def dense_decode_unsupported_reason( __all__ = [ "dense_decode_unsupported_reason", - "minimax_m3_trtllm_gen_dense_attention", - "minimax_m3_trtllm_gen_dense_context", "minimax_m3_trtllm_gen_dense_decode", "subpage_block_table", - "uniform_dense_subpage_geometry", + "uniform_dense_subpages_per_slot", "uniform_subpages_per_slot", "write_subpage_block_table", ] diff --git a/tensorrt_llm/_torch/compilation/multi_stream/auto_multi_stream.py b/tensorrt_llm/_torch/compilation/multi_stream/auto_multi_stream.py index cfc3ef93bec0..42f9bf1447ac 100644 --- a/tensorrt_llm/_torch/compilation/multi_stream/auto_multi_stream.py +++ b/tensorrt_llm/_torch/compilation/multi_stream/auto_multi_stream.py @@ -1,3 +1,17 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import time from dataclasses import dataclass, field from operator import getitem @@ -194,6 +208,20 @@ def flatten_args(args): elif isinstance(arg, torch.fx.Node) and arg.op != "placeholder": in_edges[arg] = self.nodes[arg] + if node.op == "output": + # An in-place op may mutate a graph input without returning a + # value (Eagle3 captures hidden states into a preallocated + # buffer with inplace_slice_copy), so the FX output does not + # reach that side effect. Make graph exit depend on the last + # mutation of every touched tensor: the scheduled graph then + # emits the mutation before `output` (a node emitted after + # `output` is dead code once the module is recompiled), and + # with live auxiliary streams the exit waits on the mutating + # stream before a graph-external consumer reads the buffer. + for mutated_arg, mutator in latest_inplace_stat.items(): + if isinstance(mutated_arg, torch.fx.Node): + in_edges[mutated_arg] = mutator + # For node without in edge, connect it to the entry if len(in_edges) == 0: in_edges[None] = self.entry_node diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index 1e6b481ece93..69b574c3d0e0 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -2848,9 +2848,9 @@ def __init__(self, model_config: "ModelConfig[PretrainedConfig]"): model_config = get_text_model_config(model_config) if model_config.quant_config.kv_cache_quant_algo == QuantAlgo.NVFP4: # M3's 57 sparse target layers have an MSA NVFP4 consumer, but the - # one-model Eagle layer has no shipped P32 NVFP4 decode cubin. - # Keep its modules and separate cache on their established FP8 - # representation while the target remains NVFP4. + # one-model Eagle layer has no matching NVFP4 decode cubin. Keep its + # modules and its shared draft KV layer on FP8 while the target + # remains NVFP4. model_config.extra_attrs["draft_kv_cache_quant_algo_override"] = QuantAlgo.FP8 super().__init__(MiniMaxM3Model(model_config), model_config) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index b8e160cf7cd6..3c837679fa77 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1147,11 +1147,8 @@ def _should_create_separate_draft_kv_cache(self) -> bool: self._kv_cache_manager_cls, 'supports_shared_draft_layers', True): # Under attention DP, draft layers share the target manager (the - # layout existing deployments were validated with). A manager can - # opt out: MiniMax-M3's coalesces an index-K pool into its KV - # pages and exposes only synthetic AttentionOp tensors, which the - # dense Eagle3 drafter cannot attend against, so it requires the - # separate draft manager even under attention DP. + # layout existing deployments were validated with). Managers that + # cannot expose a draft-compatible view can opt out. logger.info("Attention DP: draft layers share the target KV " "cache manager.") return False @@ -1276,23 +1273,12 @@ def _create_one_model_draft_kv_cache_manager( # the sparse_attention_config. Get it from effective_draft_config which # falls back to the target model's config for MTP mode. sparse_attn_config = effective_draft_config.sparse_attention_config - # A target manager class may request a different page size for the - # separate draft manager (e.g. MiniMax-M3, see - # draft_manager_tokens_per_block there for the rationale). - draft_tpb = getattr(self._kv_cache_manager_cls, - 'draft_manager_tokens_per_block', - self._tokens_per_block) - if draft_tpb != self._tokens_per_block: - logger.info( - f"Draft KV cache manager uses tokens_per_block={draft_tpb} " - f"(target uses {self._tokens_per_block}).") - draft_kv_config.tokens_per_block = draft_tpb return _create_kv_cache_manager( model_engine=None, kv_cache_manager_cls=draft_kv_cache_manager_cls, mapping=self._mapping, kv_cache_config=draft_kv_config, - tokens_per_block=draft_tpb, + tokens_per_block=self._tokens_per_block, max_seq_len=self._max_seq_len, max_batch_size=self._max_batch_size, spec_config=self._speculative_config, @@ -2044,8 +2030,8 @@ def _create_kv_cache_manager( # One-model spec with shared draft layers appends the drafter's # layers to this manager; tell the manager how many. Anchor on the # pretrained TARGET layer count — local num_hidden_layers may already - # include the draft tail. Consumed by managers with a draft sub-page - # view (MiniMax-M3); others ignore it. Masked/cross flows yield a + # include the draft tail. Consumed by managers with a shared draft + # view; others ignore it. Masked/cross flows yield a # non-positive delta and correctly report 0. target_num_layers = getattr(config, "num_hidden_layers", None) num_appended_draft_layers = (len(per_layer_num_kv_heads) - diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 79e28b05abe6..539f402d5bb6 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -975,13 +975,8 @@ def moe_load_balancer_iter_info(self, value: Tuple[bool, bool]): def use_beam_search(self): return self.max_beam_width > 1 - def _get_draft_kv_cache_manager( - self, resource_manager: ResourceManager - ) -> Optional[Union[KVCacheManager, KVCacheManagerV2]]: - """ - Returns the draft KV cache manager only in one-model speculative decoding - mode where the target model manages a separate draft KV cache. - """ + def _get_draft_kv_cache_manager(self, resource_manager: ResourceManager): + """Return the one-model draft manager or shared-layout adapter.""" return get_draft_kv_cache_manager(self.spec_config, resource_manager) @contextmanager diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index bf91c4ab2ff5..2af0d97de7be 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -5734,11 +5734,6 @@ def _pad_attention_dp_dummy_request(self): and self.max_num_tokens is not None): token_nums = [self.max_num_tokens] - # A separate draft KV cache manager must also see the dummy, or - # its prepare_resources hits an unknown request id. - draft_kv_cache_manager = self.resource_manager.get_resource_manager( - ResourceManagerType.DRAFT_KV_CACHE_MANAGER) - if (not self._enable_dsv4_adp_dummy_fixes or self.kv_cache_transceiver is None): try: @@ -5748,12 +5743,11 @@ def _pad_attention_dp_dummy_request(self): is_gen=self._adp_dummy_is_gen, prepare_resource=True, max_num_draft_tokens=self.max_total_draft_tokens, - draft_kv_cache_manager=draft_kv_cache_manager, ) except OutOfPagesError: dummy_requests = None if not dummy_requests: - # Both KV cache managers report allocation failure by returning + # The cache manager reports allocation failure by returning # None, expecting the caller to retry on a later iteration. An # empty batch is safe here because _can_queue() allgathers batch # sizes, so every rank skips the forward pass together. @@ -5786,7 +5780,6 @@ def _pad_attention_dp_dummy_request(self): is_gen=self._adp_dummy_is_gen, prepare_resource=True, max_num_draft_tokens=self.max_total_draft_tokens, - draft_kv_cache_manager=draft_kv_cache_manager, ) except OutOfPagesError: dummy_requests = None diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index e28096ceaf17..56250b745f97 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -511,13 +511,12 @@ def create_py_executor( if hasattr(spec_config, '_max_batch_size'): spec_config._max_batch_size = max_batch_size - # WAR for https://nvbugs/5807902 - # Disable separate draft KV cache in disaggregated mode - # Enable separate pool for None DI + Non-KVBM and Aggregated + KVBM - # (The shared-manager fallback is exactly what MiniMax-M3 wants here - # — its drafter shares the target manager and rides prefix reuse and - # KV transfer natively — so M3's earlier #17341 exemption is retired.) - if cache_transceiver_config is not None: + # WAR for https://nvbugs/5807902: disable the separate draft KV cache + # in disaggregated mode. MiniMax-M3 also uses the unified target cache + # in every supported one-model configuration; its native P128 view maps + # the dense draft K/V pages inside M3's non-uniform mega-slot. + if (cache_transceiver_config is not None + or is_minimax_m3(m3_sparse_config)): spec_config._allow_separate_draft_kv_cache = False # chunk_unit_size may be changed to 64 when using flash mla diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index df8aebdb7ff8..bea3730ea319 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -580,10 +580,8 @@ def maybe_capture_hidden_states( f"hidden_states={hidden_states.shape[0]}") to_save = hidden_states[:num_tokens] if residual is not None: - # residual shares its leading (token) dim with hidden_states - # (both come from the same decoder layer), so the bound - # check above already guarantees num_tokens <= - # residual.shape[0]; no separate check is needed. + # Both values come from the same decoder layer, so the + # hidden-state bound above also covers residual. to_save = to_save + residual[:num_tokens] inplace_slice_copy(self.hidden_states, to_save, i * self.hidden_size, @@ -794,7 +792,7 @@ def _forward_impl(self, # state the target model expects. original_all_rank_num_tokens = attn_metadata.all_rank_num_tokens - # Get the draft KV cache manager if using separate layouts + # Resolve a separate draft manager or shared-layout adapter, if needed. draft_kv_cache_manager = self.get_draft_kv_cache_manager( resource_manager) diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 7b76d8eadb13..cb05cda90016 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -122,9 +122,10 @@ def should_use_separate_draft_kv_cache(spec_config) -> bool: def prepare_attn_metadata_for_draft_replay(attn_metadata, draft_kv_cache_manager): """ - Prepare attention metadata for CUDA graph replay when using separate draft KV cache. - Swaps to draft manager and (for DSA) re-prepares indexer slot mappings for the current - batch. Call restore_attn_metadata_after_draft_replay after replay in a finally block. + Prepare attention metadata for CUDA graph replay with a distinct + draft-side KV layout. Swaps to the resolved draft manager or view and (for + DSA) re-prepares indexer slot mappings for the current batch. Call + restore_attn_metadata_after_draft_replay after replay in a finally block. Returns saved state or None if no-op. """ if draft_kv_cache_manager is None: @@ -2107,30 +2108,28 @@ def _prepare_context_input_ids(self, input_ids, num_ctx_tokens, gather_ids, return torch.empty(0, dtype=torch.int32, device="cuda") def get_draft_kv_cache_manager(self, resource_manager): - """ - Get the draft-side KV manager (separate manager or the target - manager's draft sub-page view); see resolve_draft_kv_cache_manager. - """ - if resource_manager is None: - return None - from .utils import resolve_draft_kv_cache_manager - return resolve_draft_kv_cache_manager(resource_manager) + """Get the draft-side manager or shared-layout adapter, if any.""" + from .utils import get_draft_kv_cache_manager + + return get_draft_kv_cache_manager(self.spec_config, resource_manager) @contextmanager def draft_kv_cache_context(self, attn_metadata, draft_kv_cache_manager): """ - Context manager to temporarily switch to draft KV cache manager in one-engine speculative decoding. + Temporarily switch to a distinct draft-side KV layout in one-engine + speculative decoding. This swaps both the kv_cache_manager reference AND the block offset tensors, since the target and draft KV caches have different block layouts. """ - # draft_kv_cache_manager is None if using two-engine speculative decoding or not enabling separate draft KV cache. + # None means two-engine decoding or a unified draft layout that needs + # no adapter. if draft_kv_cache_manager is None: yield return - # Only TrtllmAttentionMetadata supports separate draft KV cache layouts + # Only TrtllmAttentionMetadata supports a distinct draft KV layout. if not isinstance(attn_metadata, TrtllmAttentionMetadata): yield return @@ -2148,7 +2147,7 @@ def draft_kv_cache_context(self, attn_metadata, draft_kv_cache_manager): target_kv_cache_block_offsets = attn_metadata.kv_cache_block_offsets target_host_kv_cache_block_offsets = attn_metadata.host_kv_cache_block_offsets - # Switch to draft KV cache manager and its block offsets + # Switch to the draft-side manager or view and its block offsets. attn_metadata.kv_cache_manager = draft_kv_cache_manager attn_metadata.kv_cache_block_offsets = attn_metadata.draft_kv_cache_block_offsets attn_metadata.host_kv_cache_block_offsets = draft_kv_cache_manager.host_kv_cache_block_offsets diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index fb5956ef10a3..56ad510c1fae 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -523,43 +523,32 @@ def get_num_extra_kv_tokens(spec_config): return 0 -def resolve_draft_kv_cache_manager(resource_manager): - """Resolve the draft-side KV manager for one-model speculative decoding. +def get_draft_kv_cache_manager(spec_config, resource_manager): + """Return the draft-side manager for one-model speculative decoding. - The registered separate manager is the ground truth when present; - otherwise ask the target manager for its draft sub-page view (managers - without one — e.g. same-geometry shared drafters — resolve to None and - the drafter attends the shared manager directly). + A registered separate manager takes precedence. Otherwise, a target + manager with a non-uniform shared layout may provide a draft-side view; + uniform shared layouts need no adapter and return ``None``. """ + if (spec_config is None or resource_manager is None + or not spec_config.spec_dec_mode.use_one_engine()): + return None + from ..pyexecutor.resource_manager import ResourceManagerType draft_manager = resource_manager.get_resource_manager( ResourceManagerType.DRAFT_KV_CACHE_MANAGER) if draft_manager is not None: return draft_manager + target_manager = resource_manager.get_resource_manager( ResourceManagerType.KV_CACHE_MANAGER) - # getattr fetches the method without executing it, so a failure inside - # view construction propagates from the call itself instead of being - # silently swallowed into "no view". - get_view = getattr(target_manager, "get_draft_subpage_view", None) + # Fetch the method before calling it so construction failures propagate + # instead of being mistaken for a manager without a view. + get_view = getattr(target_manager, "get_draft_kv_cache_view", None) return get_view() if get_view is not None else None -def get_draft_kv_cache_manager(spec_config, resource_manager): - """ - Returns the draft KV cache manager only in one-model speculative decoding - mode: the separate manager when the target manages one, or the target - manager's draft sub-page view when shared draft layers run at a smaller - kernel page size (e.g. MiniMax-M3). See resolve_draft_kv_cache_manager. - """ - if spec_config is None: - return None - if not spec_config.spec_dec_mode.use_one_engine(): - return None - return resolve_draft_kv_cache_manager(resource_manager) - - def update_spec_config_from_model_config(spec_config, model_config): from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig if not isinstance(spec_config, MTPDecodingConfig): diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 9b3ff586667c..caf983978ce0 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -16,6 +16,7 @@ import json import os import sys +import time from unittest import mock import pytest @@ -7875,10 +7876,9 @@ def _run_nvfp4_eagle3_disagg(self, (attention-DP generation), False the TEP production-candidate shape. Values that are workload-tuned (1M-context sizes) or GPU-generation- tuned (memory fractions) are CI-adjusted and called out inline. - Accuracy is asserted through the router; the drafter's KV rides the - shared logical blocks, so a corrupted or dropped drafter cache - collapses accuracy. Acceptance stats stay with the aggregated arm, - which exercises the same view code. + Accuracy is asserted through the router. A chat-GSM8K acceptance probe + additionally guards the drafter KV that rides the shared logical + blocks, since accuracy alone is insensitive to rejected draft tokens. """ if not (overlap_scheduler and cuda_graph and use_msa): pytest.skip("the disagg arm pins the production serving shape " @@ -7887,7 +7887,7 @@ def _run_nvfp4_eagle3_disagg(self, speculative_config = { "decoding_type": "Eagle3", "max_draft_len": max_draft_len, - "speculative_model": f"{llm_models_root()}/MiniMax-M3-EAGLE3", + "speculative_model": f"{llm_models_root()}/MiniMax-M3-EAGLE3-GQA", "eagle3_one_model": True, } common_config = { @@ -7975,6 +7975,13 @@ def _run_nvfp4_eagle3_disagg(self, "num_postprocess_workers": 4, "stream_interval": 100, "enable_iter_perf_stats": True, + # Preserve the complete acceptance-probe window. The default + # engine and OpenAI-server buffers retain only their latest 1000 + # iterations, which samples whichever requests finish last. Full + # 200-prompt runs produced at most 2966 records, so 5000 keeps the + # complete window without making either buffer unbounded. + "max_stats_len": 5000, + "iter_stats_max_iterations": 5000, "cuda_graph_config": { "enable_padding": True, "batch_sizes": [1, 2, 4, 8, 16], @@ -8024,14 +8031,16 @@ def _run_nvfp4_eagle3_disagg(self, if inferencemax else GSM8K(model_name)) task.evaluate(llm) - # Chat-format acceptance probe — the same workload and - # thresholds as the aggregated arm's (200 GSM8K questions, - # chat template, greedy, 512 tokens; drafter card reference - # rate 0.839 / length 3.518). The disagg fixture has no - # get_stats, so the probe window comes from the generation - # worker's /metrics buffer (enable_iter_perf_stats), drained - # right before the probe; the dataset loads from the models - # root so the probe works offline like the eval does. + # Chat-format acceptance probe — 200 GSM8K questions, chat + # template, greedy, 512 tokens (drafter-card reference rate 0.839 / + # length 3.518). Four native-P128 unified-cache validation runs + # measured 0.803-0.823 / 3.410-3.468. The floors retain run-to-run + # headroom while rejecting the earlier corrupted-draft-KV result + # (0.779-0.780 / 3.338-3.339). + # The disagg fixture has no get_stats, so read the generation + # worker's bounded iteration-stat window after draining the + # preceding eval. The dataset loads from the models root so the + # probe works offline like the eval does. import requests info = requests.get(f"{llm.router_url}/cluster_info", timeout=30).json() @@ -8050,15 +8059,22 @@ def _run_nvfp4_eagle3_disagg(self, add_generation_prompt=True) for q in questions ] - requests.get(f"http://{gen_url}/metrics", timeout=30) # drain + metrics_url = f"http://{gen_url}/metrics" + # A completed request wakes a background stats collector whose + # queue-drain timeout is 0.5 s. Let it finish before clearing or + # reading the HTTP snapshot buffer. + time.sleep(1) + requests.get(metrics_url, timeout=120).raise_for_status() probe_params = SamplingParams(max_tokens=512, temperature=0) for future in [ llm.generate_async(prompt, probe_params) for prompt in chat_prompts ]: future.result() - records = requests.get(f"http://{gen_url}/metrics", - timeout=30).json() + time.sleep(1) + response = requests.get(metrics_url, timeout=120) + response.raise_for_status() + records = response.json() drafted = accepted = steps = 0 for record in records: stats = record.get("specDecodingStats") or {} @@ -8071,13 +8087,13 @@ def _run_nvfp4_eagle3_disagg(self, print(f"MiniMax-M3 Eagle3 disagg chat-GSM8K acceptance: rate=" f"{chat_rate:.3f}, mean acceptance length=" f"{chat_length:.3f} ({steps} spec iterations)") - assert chat_rate > 0.78, \ + assert chat_rate > 0.80, \ f"Eagle3 chat-GSM8K acceptance rate too low: " \ - f"{chat_rate:.3f} (threshold 0.78, reference 0.839 from " \ + f"{chat_rate:.3f} (threshold 0.80, reference 0.839 from " \ f"the drafter card)" - assert chat_length > 3.3, \ + assert chat_length > 3.4, \ f"Eagle3 chat-GSM8K acceptance length too low: " \ - f"{chat_length:.3f} (threshold 3.3, reference 3.518 from " \ + f"{chat_length:.3f} (threshold 3.4, reference 3.518 from " \ f"the drafter card)" @pytest.mark.skip_less_device(6) @@ -8137,7 +8153,7 @@ def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, return spec_config = Eagle3DecodingConfig( max_draft_len=max_draft_len, - speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3", + speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3-GQA", ) # The runtime forces tokens_per_block per implementation (128 MSA / 32 # reference). @@ -8200,9 +8216,9 @@ def drain_spec_stats(llm): task = GSM8K(model_name) task.evaluate(llm) - # Chat-format acceptance — the drafter's training distribution - # (Inferact/MiniMax-M3-EAGLE3 card: 0.839 / 3.518). Reuses the - # live engine and the cached dataset; ~20 s under CUDA graphs. + # Chat-format acceptance — the drafter's training distribution. + # Reuses the live engine and the cached dataset; ~20 s under + # CUDA graphs. questions = [ r["question"] for r in load_dataset("gsm8k", "main", split="test") @@ -8223,20 +8239,26 @@ def drain_spec_stats(llm): assert steps > 0, "no speculative iterations recorded" chat_rate = accepted / drafted chat_length = 1 + accepted / steps - # Published reference (Inferact/MiniMax-M3-EAGLE3 drafter card): - # rate 0.839, length 3.518. Our testing thresholds: rate > 0.78, - # length > 3.3. + # The GQA drafter card publishes no GSM8K figure, so the + # reference stays the MHA card's (rate 0.839, length 3.518): + # the GQA head is retrained on the same data with only the + # attention changed (64 -> 4 KV heads), the cards agree on the + # benchmark they share (MT-Bench 2.698 vs 2.668), and this arm + # measures the GQA head at 0.838 / 3.515 — indistinguishable + # from the MHA reference. Floors: rate > 0.80, length > 3.4 + # (see the disaggregated arm for how they are calibrated). ref_rate, ref_length = 0.839, 3.518 - ref_source = "the Inferact/MiniMax-M3-EAGLE3 drafter model card" + ref_source = ("the Inferact/MiniMax-M3-EAGLE3 (MHA) drafter " + "card, which the GQA head matches on GSM8K") print(f"MiniMax-M3 Eagle3 chat-GSM8K acceptance: rate=" f"{chat_rate:.3f}, mean acceptance length=" f"{chat_length:.3f} ({steps} spec iterations)") - assert chat_rate > 0.78, \ + assert chat_rate > 0.80, \ f"Eagle3 chat-GSM8K acceptance rate too low: {chat_rate:.3f} " \ - f"(threshold 0.78, reference {ref_rate} from {ref_source})" - assert chat_length > 3.3, \ + f"(threshold 0.80, reference {ref_rate} from {ref_source})" + assert chat_length > 3.4, \ f"Eagle3 chat-GSM8K acceptance length too low: " \ - f"{chat_length:.3f} (threshold 3.3, reference {ref_length} " \ + f"{chat_length:.3f} (threshold 3.4, reference {ref_length} " \ f"from {ref_source})" print(f"Eagle3 acceptance references: rate {ref_rate}, length " f"{ref_length} — from {ref_source}.") @@ -8244,7 +8266,7 @@ def drain_spec_stats(llm): @pytest.mark.skip_less_device(4) @pytest.mark.skip_less_device_memory(140000) def test_nvfp4_kv_eagle3_smoke(self): - """Generate through MSA with a true NVFP4 KV cache and linear Eagle3.""" + """Generate through MSA with a true NVFP4 KV cache and GQA Eagle3.""" from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import \ msa_package_available @@ -8254,7 +8276,7 @@ def test_nvfp4_kv_eagle3_smoke(self): model_path = f"{llm_models_root()}/MiniMax-M3-NVFP4" spec_config = Eagle3DecodingConfig( max_draft_len=3, - speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3", + speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3-GQA", ) kv_cache_config = KvCacheConfig( dtype="nvfp4", diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py b/tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py index c10b80e0e816..f718afd86558 100644 --- a/tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py +++ b/tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py @@ -158,15 +158,14 @@ def test_hybrid_target_dense_pool_stays_fp8_p128_while_sparse_is_nvfp4_p128(): try: buffers = manager.get_buffers(0, "HND") k, v = buffers[:, 0], buffers[:, 1] - data_pool, slot_stride, pages_per_role = manager.get_dense_kv_subpage_pool(0) - assert pages_per_role == 1 + data_pool, slot_stride = manager.get_kv_subpage_pool(0, "HND") assert data_pool.dtype == torch.float8_e4m3fn assert data_pool.shape[2:] == (128, HEAD_DIM) for slot in (0, int(k.shape[0]) - 1): assert data_pool[slot * slot_stride].data_ptr() == k[slot].data_ptr() assert data_pool[slot * slot_stride + 1].data_ptr() == v[slot].data_ptr() - with pytest.raises(RuntimeError, match="has no NVFP4 block-scale pool"): - manager.get_dense_kv_scale_subpage_pool(0) + with pytest.raises(RuntimeError, match="require an NVFP4 KV cache"): + manager.get_block_scale_buffers(0, "HND") dense_buffers = manager.kv_cache_manager_py_config.layers[0].buffers assert {buffer.role for buffer in dense_buffers} == { @@ -183,7 +182,7 @@ def test_hybrid_target_dense_pool_stays_fp8_p128_while_sparse_is_nvfp4_p128(): manager.shutdown() -def test_hybrid_shared_eagle_view_points_at_the_fp8_p32_pool(): +def test_hybrid_shared_eagle_view_points_at_the_fp8_p128_pool(): from tensorrt_llm.bindings import DataType manager = _create_manager( @@ -196,12 +195,12 @@ def test_hybrid_shared_eagle_view_points_at_the_fp8_p32_pool(): ) try: draft_layer = 4 - view = manager.get_draft_subpage_view() + view = manager.get_draft_kv_cache_view() assert view is not None - assert view.tokens_per_block == 32 + assert view.tokens_per_block == 128 assert view.dtype == DataType.FP8 - k, _v = manager.get_fp8_dense_buffers(draft_layer) - assert int(view.kv_cache_pool_pointers[0, 0]) == k.data_ptr() + flat_pool, _slot_stride = manager.get_kv_subpage_pool(draft_layer, "HND") + assert int(view.kv_cache_pool_pointers[0, 0]) == flat_pool.data_ptr() assert int(view.kv_cache_pool_pointers[0, 1]) == 0 assert view._source_pool_id == int( manager.kv_cache_pool_mapping[manager.layer_offsets[draft_layer], 0] @@ -211,7 +210,7 @@ def test_hybrid_shared_eagle_view_points_at_the_fp8_p32_pool(): manager.layer_offsets[draft_layer] ].buffers assert len(draft_buffers) == 2 - assert all(buffer.tokens_per_block_override == 32 for buffer in draft_buffers) + assert all(buffer.tokens_per_block_override is None for buffer in draft_buffers) finally: manager.shutdown() @@ -231,15 +230,6 @@ def test_subpage_block_table_splits_k_and_v_rows(): assert table[:, 1].tolist() == [[1, 28, 64], [19, 46, 100]] -def test_subpage_block_table_expands_logical_p128_to_physical_p32(): - slots = torch.tensor([[0, 3]], device="cuda", dtype=torch.int32) - table = subpage_block_table(slots, subpages_per_slot=36, pages_per_role=4) - - assert table.shape == (1, 2, 8) - assert table[0, 0].tolist() == [0, 1, 2, 3, 108, 109, 110, 111] - assert table[0, 1].tolist() == [4, 5, 6, 7, 112, 113, 114, 115] - - def test_subpage_block_table_reuses_one_buffer(): """All dense layers share the arena block, so a later call must not alias a live earlier one within a step; they are written before every use.""" diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py b/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py index 15eb68157018..afef7bb6cf18 100644 --- a/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py +++ b/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py @@ -9,7 +9,6 @@ integration accuracy test. """ -from inspect import signature from types import SimpleNamespace import pytest @@ -21,7 +20,6 @@ from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_scatter import ( fused_write_layer_caches, fused_write_layer_caches_nvfp4, - fused_write_subpaged_layer_caches, ) from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import ( msa_ported_decode_active, @@ -1279,63 +1277,6 @@ def fake_nvfp4(q, k, v, scale_buffers, indexes, meta, **kwargs): ) == first_ptrs -def test_fp8_subpaged_dispatch_forwards_dequant_scale( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """The shared Eagle cache path must preserve its FP8 dequant scale.""" - import tensorrt_llm._torch.attention_backend.fmha.msa_sparse_gqa as msa_gqa - import tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils as msa_utils - import tensorrt_llm._torch.attention_backend.sparse.minimax_m3.trtllm_gen_dense_decode as dense_decode - - monkeypatch.setattr( - msa_utils, "msa_decode_span_bounds", lambda metadata, num_tokens: (0, 0, 0, 0, 0) - ) - parameter = signature(dense_decode.minimax_m3_trtllm_gen_dense_attention).parameters[ - "kv_scale_quant_orig" - ] - assert parameter.default is None - captured: dict[str, object] = {} - - def fake_dense_attention(*_args: object, **kwargs: object) -> None: - captured.update(kwargs) - - monkeypatch.setattr( - dense_decode, - "minimax_m3_trtllm_gen_dense_attention", - fake_dense_attention, - ) - - attention = MiniMaxM3MsaSparseAttention.__new__(MiniMaxM3MsaSparseAttention) - attention.layer_idx = 3 - attention.head_dim = 128 - attention.num_heads = 8 - attention.q_scaling = 1.0 - metadata = SimpleNamespace( - kv_cache_manager=SimpleNamespace( - is_nvfp4_layer=lambda layer_idx: False, - is_fp8_subpaged_layer=lambda layer_idx: layer_idx == 3, - ), - _msa_prewritten_layer=3, - ) - q = torch.zeros(2, attention.num_heads * attention.head_dim) - output = torch.empty_like(q) - dequant_scale = torch.ones(3, dtype=torch.float32) - - msa_gqa.run_msa_paged_gqa( - attention, - q, - None, - None, - metadata, - output, - kv_block_indexes=None, - plan=None, - kv_scale_quant_orig=dequant_scale, - ) - - assert captured["kv_scale_quant_orig"] is dequant_scale - - @pytest.mark.parametrize("q_heads,kv_heads", [(64, 4), (16, 1)]) def test_nvfp4_standard_stage_uses_preplanned_msa_and_stable_scratch( monkeypatch, q_heads, kv_heads @@ -1662,37 +1603,6 @@ def test_fused_scatter_matches_reference(cache_dtype, num_kv_heads, with_idx): torch.testing.assert_close(idx_pool, ref_idx_pool) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -def test_fp8_subpage_scatter_places_tokens_inside_p32_pages(): - torch.manual_seed(17) - num_slots, pages_per_role, num_heads = 3, 4, 2 - physical_page, head_dim = 32, 128 - k_cache = torch.zeros( - (num_slots, pages_per_role, num_heads, physical_page, head_dim), - dtype=torch.float8_e4m3fn, - device="cuda", - ) - v_cache = torch.zeros_like(k_cache) - slots = torch.tensor([0, 31, 32, 127, 128, 255, -1], dtype=torch.int32, device="cuda") - qkv = torch.randn( - slots.numel(), 2 * num_heads * head_dim + 13, dtype=torch.bfloat16, device="cuda" - ) - k = qkv[:, : num_heads * head_dim] - v = qkv[:, num_heads * head_dim : 2 * num_heads * head_dim] - - assert fused_write_subpaged_layer_caches(k_cache, v_cache, slots, k, v) - expected_k = k.reshape(slots.numel(), num_heads, head_dim).to(torch.float8_e4m3fn) - expected_v = v.reshape(slots.numel(), num_heads, head_dim).to(torch.float8_e4m3fn) - for row, slot in enumerate(slots[:-1].tolist()): - page, logical_within = divmod(slot, 128) - subpage, within = divmod(logical_within, physical_page) - assert torch.equal(k_cache[page, subpage, :, within], expected_k[row]) - assert torch.equal(v_cache[page, subpage, :, within], expected_v[row]) - # The invalid row is masked and therefore cannot touch any cache location. - assert torch.count_nonzero(k_cache[2]).item() == 0 - assert torch.count_nonzero(v_cache[2]).item() == 0 - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def test_nvfp4_scatter_writes_physical_p32_data_and_scale_layouts(): from tensorrt_llm._utils import get_sm_version diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_unified_draft_view.py b/tests/unittest/_torch/attention/sparse/test_minimax_m3_unified_draft_view.py index 360fe6c5fb75..25a341059e7e 100644 --- a/tests/unittest/_torch/attention/sparse/test_minimax_m3_unified_draft_view.py +++ b/tests/unittest/_torch/attention/sparse/test_minimax_m3_unified_draft_view.py @@ -12,249 +12,212 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Pure-logic tests for MiniMaxM3DraftSubpageView. - -The view presents the shared manager's draft-layer pool to the attention ops -at a smaller kernel page size (32-token pages inside 128-token logical -blocks). These tests validate the addressing math against a fake manager, so -they run without GPUs: slot ``s`` of the drafter's layer must resolve to K -sub-pages ``s*scale*subdiv + j`` and V sub-pages offset by ``subdiv`` (V is -laid out immediately after K within the slot). -""" +"""Pure-logic tests for MiniMaxM3DraftKVCacheView.""" +import pytest import torch from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.cache_manager import ( - MiniMaxM3DraftSubpageView, + MiniMaxM3DraftKVCacheView, MiniMaxM3KVCacheManagerV2, derive_shared_draft_layout, ) +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm.bindings import DataType DRAFT_LAYER = 60 -SCALE = 178 # sub-pages per mega-slot, in units of the drafter's 128-tok page +SCALE = 178 # P128 pages per M3 mega-slot ADDR = 0x7000_0000 +class _FakeFlatPool: + shape = ((1024 - 1) * SCALE + 2,) + + def data_ptr(self): + return ADDR + + class _FakeManager: tokens_per_block = 128 max_blocks_per_seq = 16 num_pools = 1 - # V2's flattened bound uses the target pool's 128-token page units and - # base pointer. The draft view must not delegate this value. - blocks_in_primary_pool = 1024 * SCALE + num_attention_op_pools = 1 + enable_swa_scratch_reuse = False + dtype = DataType.FP8 + _stream = object() def __init__(self): self.layer_offsets = {DRAFT_LAYER: DRAFT_LAYER} self.kv_cache_pool_mapping = torch.zeros((DRAFT_LAYER + 1, 2), dtype=torch.int32) self.kv_cache_pool_mapping[DRAFT_LAYER] = torch.tensor([0, 7], dtype=torch.int32) - self.slot_rows = [[5, 7]] + self.kv_cache_pool_pointers = torch.tensor([[ADDR - 1024, 0]], dtype=torch.int64) + self.index_scales = torch.tensor([SCALE], dtype=torch.int32) + self.kv_offset = torch.tensor([1], dtype=torch.int32) + self.host_kv_cache_block_offsets = torch.zeros((1, 1, 2, 16), dtype=torch.int32) - def _kv_slot_geometry(self, layer_idx, kv_layout): + def get_kv_subpage_pool(self, layer_idx, kv_layout): assert layer_idx == DRAFT_LAYER - page_shape = [self.tokens_per_block, 16, 128] - return ADDR, torch.int8, 1024, SCALE, page_shape + assert kv_layout == "HND" + return _FakeFlatPool(), SCALE - def _get_batch_cache_indices_by_pool_id(self, request_ids, *, pool_id): - assert pool_id == 0 - return self.slot_rows[: len(request_ids)] + def is_fp8_dense_layer(self, layer_idx): + assert layer_idx == DRAFT_LAYER + return False class _FakeHybridManager(_FakeManager): dtype = DataType.NVFP4 - nvfp4_dense_tokens_per_block = 32 num_pools = 2 def __init__(self): super().__init__() self.kv_cache_pool_mapping[DRAFT_LAYER] = torch.tensor([1, 7], dtype=torch.int32) + self.index_scales = torch.tensor([3, SCALE], dtype=torch.int32) + self.kv_offset = torch.tensor([1, 1], dtype=torch.int32) + self.host_kv_cache_block_offsets = torch.zeros((2, 1, 2, 16), dtype=torch.int32) - def is_fp8_subpaged_layer(self, layer_idx): + def is_fp8_dense_layer(self, layer_idx): assert layer_idx == DRAFT_LAYER return True - def _fp8_dense_data_buffers(self, layer_idx): - assert layer_idx == DRAFT_LAYER - - class _Pointer: - shape = (1024, SCALE * 4, 16, 32, 128) - - @staticmethod - def data_ptr(): - return ADDR - - return _Pointer(), None, SCALE * 4, 4 - - def _get_batch_cache_indices_by_pool_id(self, request_ids, *, pool_id): - assert pool_id == 1 - return self.slot_rows[: len(request_ids)] - def _make_view(): - return MiniMaxM3DraftSubpageView(_FakeManager(), [DRAFT_LAYER], 32) + return MiniMaxM3DraftKVCacheView(_FakeManager(), [DRAFT_LAYER]) def test_view_geometry(): view = _make_view() - assert view.tokens_per_block == 32 - assert view._subdiv == 4 - assert view._slot_units == SCALE * 4 - assert view.max_blocks_per_seq == 16 * 4 + assert view.tokens_per_block == 128 + assert view.max_blocks_per_seq == 16 assert view.num_pools == view.num_attention_op_pools == 1 assert view.kv_cache_pool_pointers.tolist() == [[ADDR, 0]] - # The draft layer's mapping row is rewritten to the view's single pool. + assert view.host_kv_cache_pool_pointers.tolist() == [[ADDR, 0]] assert view.kv_cache_pool_mapping[DRAFT_LAYER].tolist() == [0, 0] - # FlashInfer wraps the pool as a flat tensor rooted at the draft K - # pointer. Its upper bound must use 32-token sub-page units and stop after - # the last slot's V pages, not delegate the target manager's 128-token - # page bound. - assert view.blocks_in_primary_pool == (1024 - 1) * SCALE * 4 + 8 - - -def test_hybrid_view_publishes_an_fp8_pool_pointer_from_the_draft_pool(): - view = MiniMaxM3DraftSubpageView(_FakeHybridManager(), [DRAFT_LAYER], 32) - expected = [[ADDR, 0]] - assert view.kv_cache_pool_pointers.tolist() == expected - assert view.host_kv_cache_pool_pointers.tolist() == expected - assert view.dtype == DataType.FP8 - assert view._source_pool_id == 1 - assert view.blocks_in_primary_pool == (1024 - 1) * SCALE * 4 + 8 + assert view.blocks_in_primary_pool == (1024 - 1) * SCALE + 2 + assert view.trtllm_gen_extra_tokens_per_block == frozenset({128}) + +def test_block_table_uses_native_p128_copy(monkeypatch): + view = _make_view() + calls = [] + + def fake_copy( + self, + dst_tensor, + request_ids, + beam_width, + num_contexts, + num_seqs, + max_blocks=None, + ): + calls.append((request_ids, beam_width, num_contexts, num_seqs, max_blocks)) + assert self.index_scales.tolist() == [SCALE] + assert self.kv_offset.tolist() == [1] + dst_tensor.zero_() + for block_idx, slot in enumerate((5, 7)): + dst_tensor[0, 0, 0, block_idx] = slot * SCALE + dst_tensor[0, 0, 1, block_idx] = slot * SCALE + 1 + + monkeypatch.setattr(KVCacheManagerV2, "copy_batch_block_offsets", fake_copy) dst = torch.full((1, 1, 2, view.max_blocks_per_seq), -7, dtype=torch.int32) - view.copy_batch_block_offsets(dst, request_ids=[123], beam_width=1, num_contexts=1, num_seqs=1) - unit = SCALE * 4 - expected_k = [5 * unit + j for j in range(4)] + [7 * unit + j for j in range(4)] - assert dst[0, 0, 0, :8].tolist() == expected_k - assert dst[0, 0, 1, :8].tolist() == [page + 4 for page in expected_k] + + view.copy_batch_block_offsets( + dst, + request_ids=[123], + beam_width=1, + num_contexts=1, + num_seqs=1, + max_blocks=9, + ) + + assert dst[0, 0, 0, :2].tolist() == [5 * SCALE, 7 * SCALE] + assert dst[0, 0, 1, :2].tolist() == [5 * SCALE + 1, 7 * SCALE + 1] + assert dst[0, 0, 0, 2:].tolist() == [0] * 14 + assert dst[0, 0, 1, 2:].tolist() == [0] * 14 + assert calls == [([123], 1, 1, 1, 9)] -def test_hybrid_view_rejects_non_p32_draft_pages(): - try: - MiniMaxM3DraftSubpageView(_FakeHybridManager(), [DRAFT_LAYER], 128) - except AssertionError as error: - assert "physical dense-cache page size P32" in str(error) - else: - raise AssertionError("expected NVFP4 Eagle draft view to require P32 pages") +def test_hybrid_view_uses_the_draft_layers_actual_fp8_pool(): + manager = _FakeHybridManager() + # A heterogeneous source pool reports its first layer's page scale, not + # the rerooted draft layer's flat-pool stride. + manager.index_scales[1] = SCALE + 11 + view = MiniMaxM3DraftKVCacheView(manager, [DRAFT_LAYER]) + + assert view.dtype == DataType.FP8 + assert view._source_pool_id == 1 + assert view.index_scales.tolist() == [SCALE] + assert view.kv_offset.tolist() == [1] + assert view.host_kv_cache_block_offsets.data_ptr() == ( + manager.host_kv_cache_block_offsets[1:2].data_ptr() + ) + assert view.kv_cache_pool_pointers.tolist() == [[ADDR, 0]] + assert view.blocks_in_primary_pool == (1024 - 1) * SCALE + 2 def test_nvfp4_manager_rejects_dynamic_tree_eagle_before_allocation(): class _DynamicTreeConfig: use_dynamic_tree = True - try: + with pytest.raises(NotImplementedError, match="block scales"): MiniMaxM3KVCacheManagerV2( dtype=DataType.NVFP4, spec_config=_DynamicTreeConfig(), ) - except NotImplementedError as error: - assert "supports linear Eagle3" in str(error) - assert "block scales" in str(error) - else: - raise AssertionError("expected NVFP4 dynamic-tree Eagle to be rejected") - - -def test_block_table_expansion(): - view = _make_view() - max_units = view.max_blocks_per_seq - dst = torch.full((1, 1, 2, max_units), -7, dtype=torch.int32) - view.copy_batch_block_offsets(dst, request_ids=[123], beam_width=1, num_contexts=1, num_seqs=1) - unit = SCALE * 4 - expect_k = [5 * unit + j for j in range(4)] + [7 * unit + j for j in range(4)] - expect_v = [v + 4 for v in expect_k] - assert dst[0, 0, 0, :8].tolist() == expect_k - assert dst[0, 0, 1, :8].tolist() == expect_v - # Unallocated tail slots clamp to slot 0, so entries tile its sub-pages — - # inert pads: kernels never read past the row's real block count (same - # property as test_bad_page_index_padding_is_safe). - assert dst[0, 0, 0, 8:].tolist() == [0, 1, 2, 3] * ((max_units - 8) // 4) - assert dst[0, 0, 1, 8:].tolist() == [4, 5, 6, 7] * ((max_units - 8) // 4) - - -def test_block_table_source_is_private_per_call(): - # The H2D copy reads its source at execution time, so every call must get - # its own staging buffer: a persistent one refilled in place would let the - # next iteration clobber a still-pending copy and the drafter would index - # another batch's blocks (nvbug 6293536). - view = _make_view() - first = view._host_block_table([[5, 7]], 1, 2, torch.int32) - second = view._host_block_table([[9, 11]], 1, 2, torch.int32) - assert first.data_ptr() != second.data_ptr() - unit = SCALE * 4 - # The first table still holds its own batch after the second call. - assert first[0, 0, :4].tolist() == [5 * unit + j for j in range(4)] - assert second[0, 0, :4].tolist() == [9 * unit + j for j in range(4)] - - -def test_bad_page_index_padding_is_safe(): - view = _make_view() - view._manager.slot_rows = [[5, -1]] - dst = torch.zeros((1, 1, 2, view.max_blocks_per_seq), dtype=torch.int32) - view.copy_batch_block_offsets(dst, request_ids=[1], beam_width=1, num_contexts=1, num_seqs=1) - # BAD_PAGE_INDEX (-1) clamps to slot 0: pad entries index pages 0..subdiv, - # never negative offsets. - assert dst[0, 0, 0, 4:8].tolist() == [0, 1, 2, 3] - assert (dst >= 0).all() def test_free_resources_is_noop(): - view = _make_view() - view.free_resources(object()) # must not raise nor touch the manager - - -def test_subdiv_one_degenerates_to_identity(): - # Retirement path (TRTLLM_M3_DRAFT_KV_TOKENS_PER_BLOCK=128): one - # sub-page per logical block, so the table is K=slot*scale, V=K+1. - view = MiniMaxM3DraftSubpageView(_FakeManager(), [DRAFT_LAYER], 128) - assert view._subdiv == 1 - assert view.tokens_per_block == 128 - assert view.max_blocks_per_seq == 16 - assert view.blocks_in_primary_pool == (1024 - 1) * SCALE + 2 - dst = torch.zeros((1, 1, 2, view.max_blocks_per_seq), dtype=torch.int32) - view.copy_batch_block_offsets(dst, request_ids=[1], beam_width=1, num_contexts=1, num_seqs=1) - assert dst[0, 0, 0, :2].tolist() == [5 * SCALE, 7 * SCALE] - assert dst[0, 0, 1, :2].tolist() == [5 * SCALE + 1, 7 * SCALE + 1] + _make_view().free_resources(object()) def test_manager_accessor_builds_and_caches_view(): - # Exercise the accessor itself (construction + the log statement), not - # just direct view construction: a stale field reference in the log - # f-string once raised AttributeError here and silently disabled the - # view. class _FakeSharedManager(_FakeManager): is_draft = False - draft_manager_tokens_per_block = 32 + sparse_layer_ids = list(range(3, 60)) def __init__(self): super().__init__() self._shared_draft_layer_ids = [DRAFT_LAYER] - self._draft_subpage_view_obj = None + self._draft_kv_cache_view_obj = None manager = _FakeSharedManager() - get_view = MiniMaxM3KVCacheManagerV2.get_draft_subpage_view + get_view = MiniMaxM3KVCacheManagerV2.get_draft_kv_cache_view view = get_view(manager) - assert isinstance(view, MiniMaxM3DraftSubpageView) - assert view.tokens_per_block == 32 - assert get_view(manager) is view # cached on second call + assert isinstance(view, MiniMaxM3DraftKVCacheView) + assert get_view(manager) is view manager_draft = _FakeSharedManager() manager_draft.is_draft = True assert get_view(manager_draft) is None -def test_view_sources_slots_from_the_draft_layers_actual_pool(): +def test_view_rejects_non_p128_manager(): + manager = _FakeManager() + manager.tokens_per_block = 32 + with pytest.raises(ValueError, match="tokens_per_block=128"): + MiniMaxM3DraftKVCacheView(manager, [DRAFT_LAYER]) + + +def test_view_rejects_multiple_draft_layers(): + with pytest.raises(ValueError, match="exactly one draft layer"): + MiniMaxM3DraftKVCacheView(_FakeManager(), [DRAFT_LAYER, DRAFT_LAYER + 1]) + + +def test_view_rejects_incompatible_source_pool_kv_offset(): manager = _FakeHybridManager() - view = MiniMaxM3DraftSubpageView(manager, [DRAFT_LAYER], 32) - assert view._source_pool_id == 1 - dst = torch.zeros((1, 1, 2, view.max_blocks_per_seq), dtype=torch.int32) - view.copy_batch_block_offsets(dst, [1], 1, 1, 1) - assert dst[0, 0, 0, :4].tolist() == [5 * SCALE * 4 + i for i in range(4)] + manager.kv_offset[1] += 1 + with pytest.raises(ValueError, match="block-table mapping is unavailable"): + MiniMaxM3DraftKVCacheView(manager, [DRAFT_LAYER]) + + +def test_view_rejects_swa_scratch_reuse(): + manager = _FakeManager() + manager.enable_swa_scratch_reuse = True + with pytest.raises(ValueError, match="SWA scratch reuse"): + MiniMaxM3DraftKVCacheView(manager, [DRAFT_LAYER]) def test_draft_layout_target_only_num_layers(): - # The M3 creation-site flow: num_layers carries the pretrained target - # count while the per-layer heads list is already draft-extended. - # Anchoring the tail on num_layers instead of the list marked target - # layer 59 as draft and dropped its index-K cache (crashed at startup). heads = [4] * 60 + [64] draft_ids, num_target = derive_shared_draft_layout(60, heads, 1) assert draft_ids == [60] @@ -262,16 +225,12 @@ def test_draft_layout_target_only_num_layers(): def test_draft_layout_equal_head_drafter(): - # The GQA Eagle head has the target's KV head count, so the heads list - # is uniform; the draft tail must still resolve from the list length - # (an equal-head drafter is invisible in the values). draft_ids, num_target = derive_shared_draft_layout(60, [4] * 61, 1) assert draft_ids == [60] assert num_target == 60 def test_draft_layout_pre_extended_num_layers(): - # Flows that pass the extended count directly must resolve identically. heads = [4] * 60 + [64] draft_ids, num_target = derive_shared_draft_layout(61, heads, 1) assert draft_ids == [60] @@ -282,7 +241,6 @@ def test_draft_layout_no_draft(): draft_ids, num_target = derive_shared_draft_layout(60, [4] * 60, 0) assert draft_ids == [] assert num_target == 60 - # Scalar heads (plain M3, no spec) fall back to num_layers. assert derive_shared_draft_layout(60, 4, 0) == ([], 60) diff --git a/tests/unittest/_torch/compilation/test_auto_multi_stream.py b/tests/unittest/_torch/compilation/test_auto_multi_stream.py new file mode 100644 index 000000000000..dac154fea5c8 --- /dev/null +++ b/tests/unittest/_torch/compilation/test_auto_multi_stream.py @@ -0,0 +1,113 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Scheduling of in-place side effects that the FX output does not reach. + +Eagle3 captures decoder hidden states into a preallocated buffer with +``inplace_slice_copy``; the drafter reads that buffer outside the compiled +graph. The multi-stream scheduler must emit every such mutation before ``output`` +(a node emitted after ``output`` is dead code once the module is recompiled) +and make the exit wait on the mutating stream. +""" + +import pytest +import torch +from torch.fx import Graph, GraphModule + +from tensorrt_llm._torch.compilation.multi_stream.auto_multi_stream import ( + MultiStreamDAG, + multi_stream_schedule, +) + +COPY = torch.ops.trtllm.inplace_slice_copy.default + + +def _capture(graph: Graph, dest, src, layer: int): + return graph.call_function( + COPY, kwargs={"dest": dest, "src": src, "dim1_start": layer, "dim1_end": layer + 1} + ) + + +def _decoder_stack(n_layers: int, capture_layers: tuple[int, ...]): + """A chain of layers; selected layers copy their output into ``dest``.""" + graph = Graph() + dest = graph.placeholder("dest") + x = graph.placeholder("x") + hidden = x + captures = [] + for layer in range(n_layers): + mm = graph.call_function(torch.ops.aten.mm.default, args=(hidden, hidden)) + add = graph.call_function(torch.ops.aten.add.Tensor, args=(mm, hidden)) + hidden = graph.call_function(torch.ops.aten.mul.Tensor, args=(add, 2.0)) + if layer in capture_layers: + captures.append(_capture(graph, dest, hidden, layer)) + out = graph.call_function(torch.ops.aten.neg.default, args=(hidden,)) + graph.output((out,)) + return GraphModule({}, graph), dest, captures + + +def _index_of(nodes, predicate): + return next(i for i, node in enumerate(nodes) if predicate(node)) + + +def test_graph_exit_depends_on_unreturned_inplace_side_effect() -> None: + graph = Graph() + dest = graph.placeholder("dest") + src = graph.placeholder("src") + returned = graph.call_function(torch.ops.aten.neg.default, args=(src,)) + mutation = _capture(graph, dest, src, 0) + output = graph.output(returned) + graph_module = GraphModule({}, graph) + + dag = MultiStreamDAG(graph_module) + assert dag.nodes[output].in_edges[dest] is dag.nodes[mutation] + + dag.assign_streams(max_num_streams=2) + scheduled = dag.create_new_graph() + scheduled.lint() + nodes = list(scheduled.nodes) + output_index = _index_of(nodes, lambda n: n.op == "output") + mutation_index = _index_of(nodes, lambda n: n.target is COPY) + assert mutation_index < output_index + + if dag.nodes[mutation].stream is not dag.nodes[output].stream: + event = dag.nodes[mutation].event + assert event is not None + assert any( + node.target is torch.ops.trtllm.wait_event and node.args == (event,) + for node in nodes[mutation_index:output_index] + ) + + +@pytest.mark.parametrize("max_num_streams", [2, 3]) +@pytest.mark.parametrize( + "n_layers,capture_layers", + [(6, (1, 3, 5)), (6, (1, 3, 4)), (12, (1, 5, 11)), (12, (1, 5, 8))], +) +def test_every_capture_precedes_output(n_layers, capture_layers, max_num_streams) -> None: + graph_module, _, captures = _decoder_stack(n_layers, capture_layers) + multi_stream_schedule(graph_module, max_num_streams) + graph_module.graph.lint() + nodes = list(graph_module.graph.nodes) + output_index = _index_of(nodes, lambda n: n.op == "output") + copy_indices = [i for i, node in enumerate(nodes) if node.target is COPY] + assert len(copy_indices) == len(captures) + assert all(i < output_index for i in copy_indices), (copy_indices, output_index) + + # Every capture must survive into the generated code ahead of `return`. + graph_module.recompile() + lines = [line.strip() for line in graph_module.code.splitlines() if line.strip()] + return_index = _index_of(lines, lambda line: line.startswith("return")) + copy_lines = [i for i, line in enumerate(lines) if "inplace_slice_copy" in line] + assert len(copy_lines) == len(captures) + assert all(i < return_index for i in copy_lines), (copy_lines, return_index) diff --git a/tests/unittest/_torch/speculative/test_draft_kv_dispatch.py b/tests/unittest/_torch/speculative/test_draft_kv_dispatch.py index 1ee84c63fab1..36b54050c07e 100644 --- a/tests/unittest/_torch/speculative/test_draft_kv_dispatch.py +++ b/tests/unittest/_torch/speculative/test_draft_kv_dispatch.py @@ -12,18 +12,16 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Dispatch regression tests for one-model draft KV resolution. +"""Dispatch regression tests for one-model draft KV resolution.""" -Both draft-KV consumers (attention-metadata setup and the drafting -loop) resolve through ``resolve_draft_kv_cache_manager``. The registered -resource is the ground truth: worker-level flags such as -``use_separate_draft_kv_cache`` can disagree with the manager-level -share decision (e.g. attention-DP sharing), which previously left the -drafting loop without a manager while metadata used one. -""" +from types import SimpleNamespace + +import pytest from tensorrt_llm._torch.pyexecutor.resource_manager import ResourceManagerType -from tensorrt_llm._torch.speculative.utils import resolve_draft_kv_cache_manager +from tensorrt_llm._torch.speculative.utils import get_draft_kv_cache_manager + +ONE_MODEL = SimpleNamespace(spec_dec_mode=SimpleNamespace(use_one_engine=lambda: True)) class _FakeResources: @@ -39,7 +37,7 @@ class _SharedTargetWithView: _view = object() - def get_draft_subpage_view(self): + def get_draft_kv_cache_view(self): return self._view @@ -55,26 +53,31 @@ def test_registered_separate_manager_wins(): ResourceManagerType.KV_CACHE_MANAGER: _SharedTargetWithView(), } ) - assert resolve_draft_kv_cache_manager(resources) is separate + assert get_draft_kv_cache_manager(ONE_MODEL, resources) is separate -def test_shared_manager_falls_back_to_subpage_view(): +def test_shared_manager_falls_back_to_draft_view(): target = _SharedTargetWithView() resources = _FakeResources({ResourceManagerType.KV_CACHE_MANAGER: target}) - assert resolve_draft_kv_cache_manager(resources) is target._view + assert get_draft_kv_cache_manager(ONE_MODEL, resources) is target._view def test_plain_shared_manager_resolves_to_none(): resources = _FakeResources({ResourceManagerType.KV_CACHE_MANAGER: _PlainSharedTarget()}) - assert resolve_draft_kv_cache_manager(resources) is None + assert get_draft_kv_cache_manager(ONE_MODEL, resources) is None def test_no_target_manager_resolves_to_none(): - assert resolve_draft_kv_cache_manager(_FakeResources({})) is None + assert get_draft_kv_cache_manager(ONE_MODEL, _FakeResources({})) is None + + +def test_missing_inputs_resolve_to_none(): + assert get_draft_kv_cache_manager(None, _FakeResources({})) is None + assert get_draft_kv_cache_manager(ONE_MODEL, None) is None class _BrokenViewTarget: - def get_draft_subpage_view(self): + def get_draft_kv_cache_view(self): raise AttributeError("max_blocks_per_seq") @@ -83,9 +86,5 @@ def test_broken_view_construction_propagates(): # downgrade to "no view": getattr fetches the bound method without # executing it, so the exception escapes from the call itself. resources = _FakeResources({ResourceManagerType.KV_CACHE_MANAGER: _BrokenViewTarget()}) - try: - resolve_draft_kv_cache_manager(resources) - except AttributeError as e: - assert "max_blocks_per_seq" in str(e) - else: - raise AssertionError("expected the construction failure to propagate") + with pytest.raises(AttributeError, match="max_blocks_per_seq"): + get_draft_kv_cache_manager(ONE_MODEL, resources)