From aed0b9ffa15068db8ffc20b4ccdbce478765f088 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Thu, 9 Jul 2026 03:56:40 +0000 Subject: [PATCH 1/5] [TRTLLM-14093][feat] Enable one-model Eagle3 speculative decoding for MiniMax-M3 Squashed port of feat/minimax-m3-eagle3 (PR #16021) onto feat/branch_m3. Target-side enablement for the Inferact/MiniMax-M3-EAGLE3 draft head on the reference (non-MSA) sparse backend: - Wire MiniMaxM3ForCausalLM as SpecDecOneEngineForCausalLM and capture Eagle3 aux hidden states at layer exit (fully TP-reduced; no cross-layer allreduce+norm fusion). - Rebase MiniMaxM3AttentionMetadata onto TrtllmAttentionMetadata (the shared per-step metadata is consumed by TRTLLM draft layers, the Eagle3 one-model worker, and engine isinstance gates; precedent: DSA and the M3 MSA metadata). - Route multi-token generation rows (spec verify: 1 + draft_len tokens) through the extend path; the decode branch stays reserved for batches where every row appends exactly one token. - Vectorized sync-free builder shared with a new on_update_kv_lens hook that re-derives seq_lens/prefix_lens/q_positions/out_cache_loc on device from the corrected kv_lens_cuda under overlap scheduler + spec (DSA pattern); in-bounds clamps cover the optimistic page-boundary overhang (max_seqlen_k SDPA width + slot gathers). - V1-family draft manager support in KVCacheManagerV2.add_dummy_requests (exception-safe); attention-DP dummy requests register in the draft manager; MiniMaxM3KVCacheManagerV2 opts out of shared draft layers under attention DP (its AttentionOp tensors are synthetic). - Creation-time guards: tree modes, disabled separate draft KV (disagg WAR), CUDA graphs, and MSA+spec (the in-builder MSA rejection is hoisted above routing so mixed batches cannot bypass it). - Accuracy test test_nvfp4_eagle3 (MMLU + GSM8K + acceptance probe, attention_dp parametrized) + reference rows. Validated on 4xB200 (NVFP4 tp4/ep4, draft_len=3, overlap scheduler on, eager) at this commit: - test_nvfp4_eagle3[attention_dp=False]: MMLU 85.50 / GSM8K 89.73 (refs 83/88), acceptance rate 0.709 / mean acceptance length 3.126 - test_nvfp4_eagle3[attention_dp=True]: MMLU 85.14 / GSM8K 91.32, acceptance rate 0.767 / mean acceptance length 3.302 - batch-1 greedy: 6.05 -> 17.05 tok/s (2.82x) - spec-off boot + generation clean (TRTLLM-Gen warmup now runs for M3 and is harmless) Signed-off-by: Zheyu Fu (cherry picked from commit 448c489e1cfaa86434b65fc69c80a38f0e36f151) Signed-off-by: Zheyu Fu --- .../sparse/minimax_m3/cache_manager.py | 6 + .../sparse/minimax_m3/triton_metadata.py | 181 +++++++++++++----- .../_torch/models/modeling_minimaxm3.py | 29 ++- tensorrt_llm/_torch/pyexecutor/_util.py | 11 +- .../_torch/pyexecutor/kv_cache_manager_v2.py | 66 +++++-- tensorrt_llm/_torch/pyexecutor/py_executor.py | 7 + .../_torch/pyexecutor/py_executor_creator.py | 44 ++++- .../defs/accuracy/references/gsm8k.yaml | 3 + .../defs/accuracy/references/mmlu.yaml | 3 + .../defs/accuracy/test_llm_api_pytorch.py | 74 +++++++ .../test_lists/qa/llm_function_core.txt | 2 + 11 files changed, 349 insertions(+), 77 deletions(-) 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 d416086fe52c..fda13446fe2b 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 @@ -150,6 +150,12 @@ class MiniMaxM3KVCacheManagerV2(KVCacheManagerV2): * ``sparse_index_dim`` — width of the index-K/V vectors. """ + # The AttentionOp-facing tensors this manager builds are synthetic + # placeholders over INDEX_KEY-coalesced pools, so one-model speculative + # draft layers must live in a separate manager even under attention DP + # (read by ``_should_create_separate_draft_kv_cache``). + supports_shared_draft_layers = False + def __init__( self, *args, diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py index be4b8ae24143..81b12cd5bce6 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py @@ -21,6 +21,7 @@ import torch from ...interface import AttentionMetadata +from ...trtllm import TrtllmAttentionMetadata from .common import build_paged_kv_slot_mapping @@ -165,7 +166,13 @@ def prepare(self) -> None: if batch_size == 0: self.max_seqlen_k = 1 else: - self.max_seqlen_k = int(self.seq_lens_cpu[:batch_size].max().item()) + max_k = int(self.seq_lens_cpu[:batch_size].max().item()) + # Optimistic overlap+spec lengths can overhang the page table at + # a page boundary (see derive_q_positions_and_cache_slots); the + # dense fallback consumes max_seqlen_k as the exact SDPA + # mask/gather width, so an unclamped overhang crashes SDPA + # (mask 385 vs K 384 on MMLU). + self.max_seqlen_k = min(max_k, int(self.req_to_token.shape[1])) def ensure_metadata_on_device( @@ -369,6 +376,53 @@ def _build_runtime_metadata_fresh( return meta, out_cache_loc +def derive_q_positions_and_cache_slots( + req_to_token: torch.Tensor, + prefix_lens: torch.Tensor, + cu_seqlens_q: torch.Tensor, + q_batch_row: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-Q-token K-side positions and KV slot ids, on device, sync-free. + + ``q_positions[t] = prefix_lens[row(t)] + (t - cu_seqlens_q[row(t)])``; + ``out_cache_loc[t] = req_to_token[row(t), q_positions[t]]``. Shared by + the metadata builder and the ``on_update_kv_lens`` re-derivation so the + two cannot drift. + """ + total_q = int(q_batch_row.shape[0]) + qbr = q_batch_row.to(torch.long) + tok = torch.arange(total_q, dtype=torch.int32, device=q_batch_row.device) + q_positions = prefix_lens[qbr] + (tok - cu_seqlens_q[qbr]) + # Overlap+spec: the first derivation runs with optimistic + # (full-acceptance) prefix_lens, which overhang the last allocated page + # when a request's true length sits at a page boundary. The overhanging + # slots are placeholders (on_update_kv_lens re-derives them from the + # corrected kv_lens before any forward consumes them), but the gather + # must stay in bounds: unclamped, the last row trips the CUDA gather + # assert and inner rows silently read the next request's slots. Clamp + # the table index only (non-inplace: ``.to`` aliases int64 inputs) — + # q_positions keep the optimistic values, whose mask width prepare() + # bounds separately. + idx = q_positions.to(torch.long).clamp(min=0, max=req_to_token.shape[1] - 1) + flat = qbr * req_to_token.shape[1] + idx + return q_positions, req_to_token.reshape(-1).index_select(0, flat) + + +def derive_decode_cache_slots(req_to_token: torch.Tensor, seq_lens: torch.Tensor) -> torch.Tensor: + """Decode-row KV slot ids (new token at position ``seq_lens[b] - 1``). + + Same in-bounds-placeholder clamp contract as + :func:`derive_q_positions_and_cache_slots`; the ``min=0`` floor + additionally covers zero-length dummy rows indexing ``-1``. + """ + rows = torch.arange(seq_lens.shape[0], device=seq_lens.device, dtype=torch.long) + idx = (seq_lens.to(torch.long) - 1).clamp_(min=0, max=req_to_token.shape[1] - 1) + flat = rows * req_to_token.shape[1] + idx + return req_to_token.reshape(-1).index_select(0, flat) + + + + def build_runtime_metadata_from_kv_manager( *, kv_cache_manager, @@ -539,21 +593,13 @@ def build_runtime_metadata_from_kv_manager( req_to_token = req_to_token_fresh slot_ids = torch.arange(batch, device=device, dtype=torch.int32) - # Compute out_cache_loc: per-new-token slot ids, in flattened order - # matching the q-token order the model layer projects. The Python - # loops below run on CPU lists derived from the CPU-resident - # ``seq_lens_cpu`` / ``prefix_lens`` / ``extend_seq_lens_cpu``, so - # no GPU sync is needed at this point. The resulting - # ``out_cache_loc`` tensor is constructed directly on ``device``. - # The ``int(...item())`` reads against ``req_to_token`` are a CPU - # sync but only ever run from ``prepare()`` (outside any CUDA-graph - # capture window) — they are not in the forward path. + # out_cache_loc must be flattened in the q-token order the model layer + # projects, or K/V lands in the wrong requests' slots. if is_prefill: if extend_seq_lens_cpu is None: raise ValueError("prefill metadata requires extend_seq_lens_cpu") if prefix_lens is None: raise ValueError("prefill metadata requires prefix_lens") - prefix_lens_cpu = prefix_lens.to("cpu").tolist() if static_buffers is not None: prefix_buf = static_buffers["prefix_lens"] prefix_src = prefix_lens.to(device=device, dtype=torch.int32, non_blocking=True) @@ -563,17 +609,18 @@ def build_runtime_metadata_from_kv_manager( prefix_lens_dev = ( prefix_lens.to(device) if prefix_lens.device != device else prefix_lens ) - out_cache_loc_list: List[int] = [] cu_q: List[int] = [0] - req_to_token_cpu = req_to_token_fresh.to("cpu") - for b in range(batch): - pref = int(prefix_lens_cpu[b]) - ext = int(extend_seq_lens_cpu[b]) - for offset in range(ext): - slot = int(req_to_token_cpu[b, pref + offset].item()) - out_cache_loc_list.append(slot) - cu_q.append(cu_q[-1] + ext) + for ext in extend_seq_lens_cpu: + cu_q.append(cu_q[-1] + int(ext)) total_q = cu_q[-1] + cu_seqlens_q_src = torch.tensor(cu_q, dtype=torch.int32, device=device) + q_batch_row_src = torch.repeat_interleave( + torch.arange(batch, device=device, dtype=torch.int32), + torch.tensor(extend_seq_lens_cpu, dtype=torch.int64, device=device), + ) + q_positions_src, out_cache_loc_src = derive_q_positions_and_cache_slots( + req_to_token, prefix_lens_dev, cu_seqlens_q_src, q_batch_row_src + ) if static_buffers is not None: if total_q > static_buffers["max_num_tokens"]: raise ValueError( @@ -581,33 +628,22 @@ def build_runtime_metadata_from_kv_manager( f"is smaller than current total_q={total_q}" ) out_cache_loc_buf = static_buffers["out_cache_loc"] - out_cache_loc_src = torch.tensor(out_cache_loc_list, dtype=torch.int32, device=device) out_cache_loc_buf[:total_q].copy_(out_cache_loc_src, non_blocking=True) out_cache_loc = out_cache_loc_buf[:total_q] cu_seqlens_q_buf = static_buffers["cu_seqlens_q"] - cu_seqlens_q_src = torch.tensor(cu_q, dtype=torch.int32, device=device) cu_seqlens_q_buf[: batch + 1].copy_(cu_seqlens_q_src, non_blocking=True) cu_seqlens_q = cu_seqlens_q_buf[: batch + 1] - # Populate persistent q_batch_row / q_positions in-place so - # the inner metadata's prepare() can leave them alone. q_batch_row_buf = static_buffers["q_batch_row"] q_positions_buf = static_buffers["q_positions"] - for b in range(batch): - start, end = cu_q[b], cu_q[b + 1] - if end > start: - q_batch_row_buf[start:end] = b - pref = int(prefix_lens_cpu[b]) - offsets = ( - torch.arange(start, end, device=device, dtype=torch.int32) - start + pref - ) - q_positions_buf[start:end].copy_(offsets, non_blocking=True) + q_batch_row_buf[:total_q].copy_(q_batch_row_src, non_blocking=True) + q_positions_buf[:total_q].copy_(q_positions_src, non_blocking=True) q_batch_row = q_batch_row_buf[:total_q] q_positions = q_positions_buf[:total_q] else: - out_cache_loc = torch.tensor(out_cache_loc_list, dtype=torch.int32, device=device) - cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device) - q_batch_row = None - q_positions = None + out_cache_loc = out_cache_loc_src + cu_seqlens_q = cu_seqlens_q_src + q_batch_row = q_batch_row_src + q_positions = q_positions_src meta = MiniMaxM3TritonSparseAttentionMetadata( is_prefill=True, req_to_token=req_to_token, @@ -622,12 +658,7 @@ def build_runtime_metadata_from_kv_manager( ) else: # Decode: the new token sits at position seq_lens[b] - 1. - seq_lens_cpu_list = seq_lens_cpu.to("cpu").tolist() - out_cache_loc_list = [] - req_to_token_cpu = req_to_token_fresh.to("cpu") - for b in range(batch): - pos = int(seq_lens_cpu_list[b]) - 1 - out_cache_loc_list.append(int(req_to_token_cpu[b, pos].item())) + out_cache_loc_src = derive_decode_cache_slots(req_to_token, seq_lens_dev) if static_buffers is not None: if batch > static_buffers["max_num_tokens"]: raise ValueError( @@ -635,11 +666,10 @@ def build_runtime_metadata_from_kv_manager( f"is smaller than current batch={batch}" ) out_cache_loc_buf = static_buffers["out_cache_loc"] - out_cache_loc_src = torch.tensor(out_cache_loc_list, dtype=torch.int32, device=device) out_cache_loc_buf[:batch].copy_(out_cache_loc_src, non_blocking=True) out_cache_loc = out_cache_loc_buf[:batch] else: - out_cache_loc = torch.tensor(out_cache_loc_list, dtype=torch.int32, device=device) + out_cache_loc = out_cache_loc_src meta = MiniMaxM3TritonSparseAttentionMetadata( is_prefill=False, req_to_token=req_to_token, @@ -651,8 +681,20 @@ def build_runtime_metadata_from_kv_manager( return meta, out_cache_loc -class MiniMaxM3AttentionMetadata(AttentionMetadata): - """:class:`AttentionMetadata` that pre-builds MiniMax-M3 metadata. +class MiniMaxM3AttentionMetadata(TrtllmAttentionMetadata): + """:class:`TrtllmAttentionMetadata` that pre-builds MiniMax-M3 metadata. + + Subclassing :class:`TrtllmAttentionMetadata` (precedent: + ``DSAtrtllmAttentionMetadata``, and the MSA-path metadata) rather than + the plain :class:`AttentionMetadata` is required for one-model + speculative decoding (Eagle3): the draft layers live in the same engine + and run :class:`TrtllmAttention` against this shared per-step metadata, + so it must carry the TRTLLM surface (``kv_lens_cuda``, + ``host_request_types``, ``kv_cache_block_offsets``, + ``draft_kv_cache_block_offsets``, ``update_spec_dec_param``, ...). The + engine also gates spec-dec plumbing (draft KV cache swap, + overlap-scheduler ``kv_lens_cuda`` fixups) on + ``isinstance(..., TrtllmAttentionMetadata)``. Overrides :meth:`prepare` so the M3-sparse :class:`MiniMaxM3TritonSparseAttentionMetadata` and the per-new-token @@ -862,11 +904,56 @@ def prepare(self) -> None: "out_cache_loc": out_cache_loc, } + def on_update_kv_lens(self) -> None: + """Re-derive the M3 attachment from the corrected ``kv_lens_cuda``. + + Under the overlap scheduler + speculative decoding, prepare() + runs with optimistic cached counts (full draft acceptance) and + the engine corrects ``kv_lens_cuda`` on device before invoking + this hook (same pattern as ``DSAtrtllmAttentionMetadata``). + On-device, sync-free, and idempotent; ``seq_lens_cpu`` / + ``max_seqlen_k`` keep the optimistic values — they only bound + arange widths that the kernels mask by ``seq_lens``. + """ + super().on_update_kv_lens() + attachment = self.minimax_m3 + if not attachment: + return + meta = attachment["metadata"] + out_cache_loc = attachment["out_cache_loc"] + batch = int(meta.slot_ids.shape[0]) + kv_lens = self.kv_lens_cuda[:batch] + meta.seq_lens[:batch].copy_(kv_lens) + if meta.is_prefill: + # Only the K-side prefix moves with rejections; the Q-side + # structure (cu_seqlens_q, q_batch_row) is fixed per step. + total_q = int(meta.q_positions.shape[0]) + cu = meta.cu_seqlens_q + meta.prefix_lens[:batch].copy_(kv_lens - (cu[1 : batch + 1] - cu[:batch])) + q_positions, cache_slots = derive_q_positions_and_cache_slots( + meta.req_to_token, + meta.prefix_lens[:batch], + cu, + meta.q_batch_row[:total_q], + ) + meta.q_positions[:total_q].copy_(q_positions) + out_cache_loc[:total_q].copy_(cache_slots) + else: + # Reached today only as an identity (the hook also fires + # pre-correction on ordinary decode steps); re-deriving + # keeps corrected 0-draft steps correct once dynamic + # draft lengths make them reachable. + out_cache_loc[:batch].copy_( + derive_decode_cache_slots(meta.req_to_token, kv_lens) + ) + __all__ = [ "MiniMaxM3AttentionMetadata", "MiniMaxM3TritonSparseAttentionMetadata", "allocate_minimax_m3_static_buffers", "build_runtime_metadata_from_kv_manager", + "derive_decode_cache_slots", + "derive_q_positions_and_cache_slots", "ensure_metadata_on_device", ] diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index fbd938a24599..461fbc981242 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -62,6 +62,7 @@ ) from ..modules.multi_stream_utils import maybe_execute_in_parallel from ..modules.rms_norm import RMSNorm +from ..speculative import SpecMetadata from ..utils import ( ActivationType, AuxStreamType, @@ -69,9 +70,9 @@ get_model_extra_attrs, is_torch_compiling, ) +from .modeling_speculative import SpecDecOneEngineForCausalLM from .modeling_utils import ( DecoderModel, - DecoderModelForCausalLM, ModelConfig, filter_weights, register_auto_model, @@ -1606,6 +1607,7 @@ def forward( hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, residual: Optional[torch.Tensor], + spec_metadata: Optional[SpecMetadata] = None, **kwargs, ) -> torch.Tensor: # NVTX markers below are emitted only when TLLM_NVTX_DEBUG=1 (or @@ -1636,6 +1638,10 @@ def forward( hidden_states = self.block_sparse_moe(hidden_states, attn_metadata) else: hidden_states = self.mlp(hidden_states) + # hidden_states is fully TP-reduced at layer exit (no cross-layer + # allreduce+norm fusion). + if spec_metadata is not None and spec_metadata.is_layer_capture(self.layer_idx): + spec_metadata.maybe_capture_hidden_states(self.layer_idx, hidden_states, residual) return hidden_states, residual @@ -1688,6 +1694,7 @@ def forward( input_ids: Optional[torch.IntTensor] = None, position_ids: Optional[torch.IntTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, + spec_metadata: Optional[SpecMetadata] = None, **kwargs, ) -> torch.Tensor: if (input_ids is None) ^ (inputs_embeds is not None): @@ -1698,6 +1705,7 @@ def forward( hidden_states = inputs_embeds residual = None +<<<<<<< HEAD for layer_idx, decoder_layer in enumerate(self.layers): # Per-layer NVTX range (layer0, layer1, ...). Emitted only when # TLLM_NVTX_DEBUG=1 (or TLLM_LLMAPI_ENABLE_NVTX=1) is set. @@ -1708,6 +1716,16 @@ def forward( attn_metadata=attn_metadata, residual=residual, ) +======= + for decoder_layer in self.layers: + hidden_states, residual = decoder_layer( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + residual=residual, + spec_metadata=spec_metadata, + ) +>>>>>>> bad119b1ce ([TRTLLM-14093][feat] Enable one-model Eagle3 speculative decoding for MiniMax-M3) hidden_states, _ = self.norm(hidden_states, residual) return hidden_states @@ -1756,19 +1774,14 @@ def _load_index_qk_proj_weights(model: nn.Module, weights) -> None: @register_auto_model("MiniMaxM3SparseForCausalLM") -class MiniMaxM3ForCausalLM(DecoderModelForCausalLM[MiniMaxM3Model, PretrainedConfig]): +class MiniMaxM3ForCausalLM(SpecDecOneEngineForCausalLM[MiniMaxM3Model, PretrainedConfig]): """Text-only M3 model.""" def __init__(self, model_config: "ModelConfig[PretrainedConfig]"): raw_pretrained = model_config.pretrained_config if is_minimax_m3_vl_config(raw_pretrained): model_config = get_text_model_config(model_config) - super().__init__( - MiniMaxM3Model(model_config), - config=model_config, - hidden_size=model_config.pretrained_config.hidden_size, - vocab_size=model_config.pretrained_config.vocab_size, - ) + super().__init__(MiniMaxM3Model(model_config), model_config) def load_weights(self, weights, *args, **kwargs): # Fuse index_q/index_k into each index_qk_proj module (also covers diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 9049b1b90a73..78794c42fda0 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1047,10 +1047,13 @@ def _should_create_separate_draft_kv_cache(self) -> bool: in the target model and don't produce a separate ModelConfig. We fall back to the target model's config via _get_effective_draft_config(). """ - if self._mapping.enable_attention_dp: - logger.info( - "Attention DP is enabled, separate draft KV cache is not supported." - ) + if self._mapping.enable_attention_dp and getattr( + self._kv_cache_manager_cls, 'supports_shared_draft_layers', + True): + # Back-compat: attention DP keeps the shared-manager layout + # existing deployments were validated with. + logger.info("Attention DP: draft layers share the target KV " + "cache manager.") return False return should_use_separate_draft_kv_cache(self._speculative_config) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index ba19d7ece768..5e0f5b6f4920 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -2783,22 +2783,39 @@ def release_resources( return None draft_kv_cache = None if draft_kv_cache_manager is not None: - draft_kv_cache = draft_kv_cache_manager._create_kv_cache( - req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy - ) - # Dummy path: see comment above, no salt. - if draft_kv_cache is None: - release_resources(req) - return None - success = draft_kv_cache.resume(draft_kv_cache_manager._stream.cuda_stream) - if not success: - release_resources(req, free_draft_resources=True) - return None - draft_kv_cache.stop_committing() - success = draft_kv_cache.resize(dummy_capacity) - if not success: - release_resources(req, free_draft_resources=True) - return None + if isinstance(draft_kv_cache_manager, KVCacheManagerV2): + draft_kv_cache = draft_kv_cache_manager._create_kv_cache( + req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy + ) + # Dummy path: see comment above, no salt. + if draft_kv_cache is None: + release_resources(req) + return None + success = draft_kv_cache.resume(draft_kv_cache_manager._stream.cuda_stream) + if not success: + release_resources(req, free_draft_resources=True) + return None + draft_kv_cache.stop_committing() + success = draft_kv_cache.resize(dummy_capacity) + if not success: + release_resources(req, free_draft_resources=True) + return None + else: + # V1-family draft manager (no per-request cache handles); + # mirrors KVCacheManager.add_dummy_requests. The C++ side + # raises on allocation failure rather than returning a + # status, so release before propagating. + draft_seq_added = False + try: + draft_kv_cache_manager.impl.add_sequence_batch( + [(req.py_request_id, token_num, beam_width)], [req] + ) + draft_seq_added = True + for _ in range(self.num_extra_kv_tokens): + draft_kv_cache_manager.impl.add_token(req.py_request_id) + except Exception: + release_resources(req, free_draft_resources=draft_seq_added) + raise if is_gen: req.state = LlmRequestState.GENERATION_IN_PROGRESS @@ -2809,13 +2826,28 @@ def release_resources( new_capacity = kv_cache.capacity + _kv_draft + 1 success = kv_cache.resize(new_capacity, history_length=history_hint) if not success: - release_resources(req, free_draft_resources=draft_kv_cache is not None) + # V1-family draft allocations have no draft_kv_cache + # handle, so key on the manager, not the handle. + release_resources( + req, + free_draft_resources=draft_kv_cache_manager is not None, + ) return None if draft_kv_cache is not None: success = draft_kv_cache.resize(new_capacity) if not success: release_resources(req, free_draft_resources=True) return None + elif draft_kv_cache_manager is not None: + # Gen dummies must expose a 1 + draft_len kv span to + # the draft layers; a V1 manager only grows a + # sequence via add_token. + try: + for _ in range(_kv_draft): + draft_kv_cache_manager.impl.add_token(req.py_request_id) + except Exception: + release_resources(req, free_draft_resources=True) + raise if use_mrope: _populate_dummy_mrope_config(req, token_num, is_gen) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 3ab3b8d4f6d9..59f2fdf218e0 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -5728,6 +5728,11 @@ 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): llm_request = self.kv_cache_manager.add_dummy_requests( @@ -5736,6 +5741,7 @@ 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, )[0] llm_request.is_attention_dp_dummy = True spec_resource_manager = self.resource_manager.get_resource_manager( @@ -5760,6 +5766,7 @@ 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 57aa45e606f6..89aa328e00ae 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -35,7 +35,8 @@ from ..attention_backend.trtllm import TrtllmAttention from ..distributed import Distributed from ..speculative import (get_num_extra_kv_tokens, get_spec_drafter, - get_spec_resource_manager) + get_spec_resource_manager, + should_use_separate_draft_kv_cache) from ..virtual_memory import scope as virtual_memory_scope from ._util import (KvCacheCreator, _adjust_torch_mem_fraction, create_py_executor_instance, instantiate_sampler, is_mla, @@ -675,6 +676,47 @@ def drafting_loop_wrapper(model): max_num_tokens = model_engine.max_num_tokens sparse_attention_config = model_engine.sparse_attention_config + # The MSA kernel path packs one query token per generation row and its + # prefill plan stages CPU values that go stale under overlap kv-len + # correction, so it cannot verify draft tokens (two-model spec also + # emits multi-token target rows) — reject at creation instead of at the + # first verify step inside prepare(). + if (sparse_attention_config is not None + and sparse_attention_config.algorithm == "minimax_m3" + and getattr(sparse_attention_config, "sparse_use_msa", False) + and spec_config is not None): + raise NotImplementedError( + "Speculative decoding is not supported with the MiniMax-M3 MSA " + "kernel path (sparse_use_msa=True): MSA packs one query token " + "per generation row. Set sparse_use_msa=False to use the " + "reference sparse backend with speculative decoding.") + + if (sparse_attention_config is not None + and sparse_attention_config.algorithm == "minimax_m3" + and spec_config is not None + and spec_config.spec_dec_mode.is_eagle3_one_model()): + if not spec_config.is_linear_tree: + raise NotImplementedError( + "Tree-based speculative decoding (eagle_choices / " + "use_dynamic_tree) is not supported with MiniMax-M3 sparse " + "attention: the M3 sparse kernels implement linear-chain " + "verification only. Remove eagle_choices / use_dynamic_tree " + "from the speculative config.") + if llm_args.cuda_graph_config is not None: + raise NotImplementedError( + "CUDA graphs are not supported with MiniMax-M3 sparse " + "attention and speculative decoding: multi-token verify " + "routes through the M3 extend path, which is not " + "capture-safe yet. Set cuda_graph_config to null.") + if not should_use_separate_draft_kv_cache(spec_config): + raise NotImplementedError( + "One-model speculative decoding with MiniMax-M3 sparse " + "attention requires a separate draft KV cache manager, but " + "it is disabled for this configuration (e.g. disaggregated " + "serving disables it as a WAR for nvbug 5807902). Use " + "two-model speculative decoding (eagle3_one_model=False) " + "instead.") + # Set default value for cache_transceiver_config.max_tokens_in_buffer if cache_transceiver_config and cache_transceiver_config.max_tokens_in_buffer is None: cache_transceiver_config.max_tokens_in_buffer = net_max_seq_len diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index e4248e9d327d..f61bef4ba921 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -485,6 +485,9 @@ nvidia/MiniMax-M3-NVFP4: - quant_algo: MIXED_PRECISION kv_cache_quant_algo: FP8 accuracy: 86 + - quant_algo: MIXED_PRECISION + spec_dec_algo: Eagle3 + accuracy: 88 nvidia/NVIDIA-Nemotron-Nano-9B-v2: - accuracy: 85.027 - quant_algo: FP8 diff --git a/tests/integration/defs/accuracy/references/mmlu.yaml b/tests/integration/defs/accuracy/references/mmlu.yaml index a870fbba6402..c332eb57e35b 100644 --- a/tests/integration/defs/accuracy/references/mmlu.yaml +++ b/tests/integration/defs/accuracy/references/mmlu.yaml @@ -287,6 +287,9 @@ nvidia/MiniMax-M3-NVFP4: - quant_algo: MIXED_PRECISION kv_cache_quant_algo: FP8 accuracy: 81 + - quant_algo: MIXED_PRECISION + spec_dec_algo: Eagle3 + accuracy: 83 moonshotai/Kimi-K2-Instruct: - quant_algo: FP8_BLOCK_SCALES accuracy: 87.65 diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 57864818f2f2..aac089c1c872 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -7664,6 +7664,80 @@ def test_nvfp4(self, use_msa): task = GSM8K(model_name) task.evaluate(llm) + @pytest.mark.skip_less_device(4) + @pytest.mark.skip_less_device_memory(140000) + @parametrize_with_ids("overlap_scheduler", [True]) + @parametrize_with_ids("attention_dp", [False, True]) + @parametrize_with_ids("tp_size,ep_size", [(4, 4)]) + def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, + overlap_scheduler): + model_name = "nvidia/MiniMax-M3-NVFP4" + model_path = f"{llm_models_root()}/MiniMax-M3-NVFP4" + max_draft_len = 3 + spec_config = Eagle3DecodingConfig( + max_draft_len=max_draft_len, + speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3", + ) + kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.6, + enable_block_reuse=False) + with LLM(model_path, + tensor_parallel_size=tp_size, + moe_expert_parallel_size=ep_size, + kv_cache_config=kv_cache_config, + sparse_attention_config=MiniMaxM3SparseAttentionConfig(), + moe_config=MoeConfig(backend="CUTLASS"), + max_seq_len=4096, + speculative_config=spec_config, + cuda_graph_config=None, + disable_overlap_scheduler=not overlap_scheduler, + enable_attention_dp=attention_dp, + trust_remote_code=True) as llm: + assert llm.args.quant_config.quant_algo == QuantAlgo.MIXED_PRECISION + task = MMLU(model_name) + task.evaluate(llm) + task = GSM8K(model_name) + task.evaluate(llm) + + # Acceptance probe (pattern: TestNemotronV3Ultra + # test_nvfp4_4gpu_mtp_ar): stream a few greedy prompts and + # derive per-step acceptance from the token increments. + raw_prompts = [ + "Solve step by step: what is 12 times 17?", + "Write a Python function that reverses a linked list.", + "The capital of France is", + ] + prompts = [ + llm.tokenizer.apply_chat_template( + [{ + "role": "user", + "content": p + }], + tokenize=False, + add_generation_prompt=True, + ) for p in raw_prompts + ] + tok_ids = [llm.tokenizer.encode(p) for p in prompts] + sampling_params = SamplingParams(max_tokens=128, temperature=0) + total_drafted = 0 + total_accepted = 0 + total_steps = 0 + for i in range(len(tok_ids)): + num_tokens = 0 + for output in llm.generate_async(tok_ids[i], + sampling_params, + streaming=True): + new_tokens = output.outputs[0].token_ids + total_drafted += max_draft_len + total_accepted += len(new_tokens) - num_tokens - 1 + total_steps += 1 + num_tokens = len(new_tokens) + accept_rate = total_accepted / total_drafted + accept_length = 1 + total_accepted / total_steps + print(f"MiniMax-M3 Eagle3 acceptance: rate={accept_rate:.3f}, " + f"mean acceptance length={accept_length:.3f}") + assert accept_rate > 0.25, \ + f"Eagle3 acceptance rate too low: {accept_rate:.3f}" + @skip_pre_blackwell class TestGLM5FP8(LlmapiAccuracyTestHarness): diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index c79b7c3f1af9..5d8e877b892d 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -689,6 +689,8 @@ accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8[use_msa=True] TIMEOU accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8_piecewise_cuda_graph[tp_size=8-ep_size=8] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=False] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=False-overlap_scheduler=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=True-overlap_scheduler=True] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_auto_dtype accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_fp8 accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] From ac89de9031165d95dc0f05a609c31d7e3667c482 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Fri, 10 Jul 2026 03:42:27 +0000 Subject: [PATCH 2/5] [TRTLLM-14093][chore] Default one-model draft KV manager to V2 under V2 targets A structurally-V2 target (sparse-attention managers like MiniMax-M3's, which don't set use_kv_cache_manager_v2) paired with a plain-transformer draft (Eagle3) resolved the draft manager to V1, requiring V1-family special cases in KVCacheManagerV2.add_dummy_requests and leaving a latent AttributeError: KVCacheV2Scheduler calls suspend_request() on the draft manager, which only exists on V2. Promote the draft class to KVCacheManagerV2 whenever the target is V2 (shared helper used by both the creation and the cache-cost estimation paths, which previously disagreed). add_dummy_requests reverts to its original V2-only shape - kv_cache_manager_v2.py returns byte-identical to its pre-enablement upstream state. The V1-draft branches were reachable by exactly one configuration - MiniMax-M3 + one-model Eagle3; every other V2-target spec config already resolves a V2 draft (flag or target-config fallback) or shares the target manager. The V2-draft combination was validated on 4xB200: MMLU 84.84 / GSM8K 90.07 (explicit flag), plus boot+acceptance probe on the new default path (AR 0.465/AL 2.40, matching the V1-draft band). Signed-off-by: Zheyu Fu (cherry picked from commit 3854f4df89c9ee03895cdef62adf7cc064fdf177) Signed-off-by: Zheyu Fu --- tensorrt_llm/_torch/pyexecutor/_util.py | 21 ++++-- .../_torch/pyexecutor/kv_cache_manager_v2.py | 66 +++++-------------- 2 files changed, 32 insertions(+), 55 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 78794c42fda0..1fee4a11808d 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -479,10 +479,8 @@ def _get_kv_size_per_token(self, # External drafter: layers start from 0, normal PP distribution # Resolve draft manager class from draft config — may differ # from target (e.g. hybrid target + plain transformer draft). - draft_kv_cache_manager_cls = get_kv_cache_manager_cls( - effective_draft_config, - kv_cache_config, - is_disagg=self._is_disagg) + draft_kv_cache_manager_cls = self._get_draft_kv_cache_manager_cls( + effective_draft_config, kv_cache_config) total += self._per_manager_cache_cost( draft_kv_cache_manager_cls, effective_draft_config, kv_cache_config) @@ -1084,6 +1082,17 @@ def _get_num_draft_layers(self) -> int: return self._draft_config.pretrained_config.num_hidden_layers return get_num_spec_layers(self._speculative_config) + def _get_draft_kv_cache_manager_cls(self, effective_draft_config, + draft_kv_config): + """Resolve the draft manager class from the draft config, promoted + to V2 when the target manager is V2.""" + draft_kv_cache_manager_cls = get_kv_cache_manager_cls( + effective_draft_config, draft_kv_config, is_disagg=self._is_disagg) + if self._is_kv_cache_manager_v2 and not issubclass( + draft_kv_cache_manager_cls, KVCacheManagerV2): + draft_kv_cache_manager_cls = KVCacheManagerV2 + return draft_kv_cache_manager_cls + def _create_one_model_draft_kv_cache_manager( self, estimating_kv_cache: bool = False, @@ -1148,8 +1157,8 @@ def _create_one_model_draft_kv_cache_manager( f"Derived draft KV cache max_attention_window for separate " f"draft manager: {draft_kv_config.max_attention_window}") # Get the appropriate KV cache manager class for the draft model - draft_kv_cache_manager_cls = get_kv_cache_manager_cls( - effective_draft_config, draft_kv_config, is_disagg=self._is_disagg) + draft_kv_cache_manager_cls = self._get_draft_kv_cache_manager_cls( + effective_draft_config, draft_kv_config) draft_kv_cache_manager_cls = self._fallback_if_unsupported_kv_cache_manager_v2( draft_kv_cache_manager_cls, effective_draft_config, draft_kv_config) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 5e0f5b6f4920..ba19d7ece768 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -2783,39 +2783,22 @@ def release_resources( return None draft_kv_cache = None if draft_kv_cache_manager is not None: - if isinstance(draft_kv_cache_manager, KVCacheManagerV2): - draft_kv_cache = draft_kv_cache_manager._create_kv_cache( - req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy - ) - # Dummy path: see comment above, no salt. - if draft_kv_cache is None: - release_resources(req) - return None - success = draft_kv_cache.resume(draft_kv_cache_manager._stream.cuda_stream) - if not success: - release_resources(req, free_draft_resources=True) - return None - draft_kv_cache.stop_committing() - success = draft_kv_cache.resize(dummy_capacity) - if not success: - release_resources(req, free_draft_resources=True) - return None - else: - # V1-family draft manager (no per-request cache handles); - # mirrors KVCacheManager.add_dummy_requests. The C++ side - # raises on allocation failure rather than returning a - # status, so release before propagating. - draft_seq_added = False - try: - draft_kv_cache_manager.impl.add_sequence_batch( - [(req.py_request_id, token_num, beam_width)], [req] - ) - draft_seq_added = True - for _ in range(self.num_extra_kv_tokens): - draft_kv_cache_manager.impl.add_token(req.py_request_id) - except Exception: - release_resources(req, free_draft_resources=draft_seq_added) - raise + draft_kv_cache = draft_kv_cache_manager._create_kv_cache( + req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy + ) + # Dummy path: see comment above, no salt. + if draft_kv_cache is None: + release_resources(req) + return None + success = draft_kv_cache.resume(draft_kv_cache_manager._stream.cuda_stream) + if not success: + release_resources(req, free_draft_resources=True) + return None + draft_kv_cache.stop_committing() + success = draft_kv_cache.resize(dummy_capacity) + if not success: + release_resources(req, free_draft_resources=True) + return None if is_gen: req.state = LlmRequestState.GENERATION_IN_PROGRESS @@ -2826,28 +2809,13 @@ def release_resources( new_capacity = kv_cache.capacity + _kv_draft + 1 success = kv_cache.resize(new_capacity, history_length=history_hint) if not success: - # V1-family draft allocations have no draft_kv_cache - # handle, so key on the manager, not the handle. - release_resources( - req, - free_draft_resources=draft_kv_cache_manager is not None, - ) + release_resources(req, free_draft_resources=draft_kv_cache is not None) return None if draft_kv_cache is not None: success = draft_kv_cache.resize(new_capacity) if not success: release_resources(req, free_draft_resources=True) return None - elif draft_kv_cache_manager is not None: - # Gen dummies must expose a 1 + draft_len kv span to - # the draft layers; a V1 manager only grows a - # sequence via add_token. - try: - for _ in range(_kv_draft): - draft_kv_cache_manager.impl.add_token(req.py_request_id) - except Exception: - release_resources(req, free_draft_resources=True) - raise if use_mrope: _populate_dummy_mrope_config(req, token_num, is_gen) From 61a5d0ed46334915ad22c2dba8a045d52c9c816c Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Fri, 10 Jul 2026 19:07:13 +0000 Subject: [PATCH 3/5] [TRTLLM-14093][feat] MSA multi-token decode for one-model Eagle3 spec verify Removes the MSA+speculative-decoding restriction: the MSA kernel path (sparse_use_msa=True) now verifies draft tokens through a multi-token decode driver instead of rejecting at creation. Driver (decode_wrapper): hybrid scheme, validated bit-exact against the eager fmha_sm100 api. The proxy (OnlyScore) pass runs natively multi-token: the kernel's inclusive causal bound with offset = kv_len - qo_len produces exactly the verify ladder (token t attends kv_len - qo_len + t + 1 positions). Top-k selection gets per-token ladder valid-page counts, so a draft token cannot select blocks past its own attend bound. The sparse GQA pass row-expands each token to a qo_len=1 pseudo-row keeping its request's full kv_len and page-table base, with the ladder in the per-row offset - the same transform the eager api applies internally; the kernel consumes kv_block_indexes per row, so native multi-token sparse is not expressible. All device ops, capture-safe. Routing/metadata: on MSA, pure-generation uniform multi-token batches stay DECODE-shaped (decode_qo_len = 1 + draft_len), keeping the captured/overlap-safe device-plan path; mixed context+gen batches take the eager extend path; the reference backend is unchanged. KV slot staging generalizes to the causal ladder, and dense layers 0-2 get a ladder SDPA mask (identical to the old math at qo_len=1). Overlap-correction fix: the MSA hook now re-stages the flat page table on mixed batches, not just the CPU length mirrors - a correction that shrinks a row across a page boundary changes the indptr layout the eager kernels rebuild from the corrected lens, misbasing every subsequent row's pages (symptom: MMLU passes, GSM8K collapses). Tests: driver bit-diff suite extended with qo_len=4 legs - proxy bit-diff vs the eager api, top-k vs a per-token ladder reference, and CUDA-graph capture/replay with mutated lens. test_nvfp4_eagle3 runs the MSA path (use_msa single-choice parametrize); QA-list rows updated to the new test IDs (the old rows no longer matched any collected test). Validation (4xB200, TP4/EP4, NVFP4): MSA adp=False MMLU 84.97 / GSM8K 90.18, AR 0.698 / AL 3.094; MSA adp=True 85.21 / 91.24, AR 0.718 / AL 3.153; reference regression control 85.04 / 90.33, AR 0.707 / AL 3.121. MSA suite wall-clock 6:10 vs reference 11:50 on the same gates. Signed-off-by: Zheyu Fu (cherry picked from commit 4d44c5897034d9ff8b5e5f2f361b480f4eaf4fbf) Signed-off-by: Zheyu Fu --- .../sparse/minimax_m3/msa_backend.py | 56 +++++++++++++------ .../sparse/minimax_m3/triton_metadata.py | 14 ++++- .../_torch/models/modeling_minimaxm3.py | 53 +++++++----------- .../_torch/pyexecutor/py_executor_creator.py | 15 ----- .../defs/accuracy/test_llm_api_pytorch.py | 14 ++++- .../test_lists/qa/llm_function_core.txt | 4 +- .../sparse/test_minimax_m3_msa_backend.py | 46 +++++++++++++++ 7 files changed, 131 insertions(+), 71 deletions(-) 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 892e154d8263..bfa81832a181 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 @@ -385,37 +385,55 @@ def _create_msa_buffers(self) -> None: ) self._alloc_msa_proxy_scratch( num_index_heads=params.num_index_heads, - max_batch=max_num_sequences, + max_tokens=self._msa_max_decode_tokens(), max_k_tiles=max_k_tiles, capture_graph=capture_graph, ) self._msa_buffers_ready = True + def _msa_max_decode_tokens(self) -> int: + """Worst-case decode-step query-token count for scratch sizing. + + Non-spec decode has one query token per sequence; one-model Eagle3 + spec verify has 1 + draft_len per gen row. max_num_tokens bounds the + step's total token count in both cases and keeps the scratch a few + MB, so it is used directly rather than plumbing the draft length. + Falls back to 0 on partially built metadata (structural tests); + callers floor the result with their own batch/token counts. + """ + max_seqs = int(getattr(self, "max_num_sequences", 0) or 0) + max_toks = int(getattr(self, "max_num_tokens", 0) or 0) + if max_toks <= 0: + return max_seqs + return max(max_seqs, min(max_toks, 16384)) + def _alloc_msa_proxy_scratch( self, *, num_index_heads: int, - max_batch: int, + max_tokens: int, max_k_tiles: int, capture_graph: bool, ) -> None: """Allocate the flat proxy max-score store and the valid-block scratch. - The store is sized for the worst-case max_k_tiles so one allocation - serves every decode step. msa_proxy_max_score_view slices the per-step - shape out of it. + The store is sized for the worst-case max_k_tiles and the worst-case + per-step query-token count (which exceeds the batch size under + speculative multi-token verify), so one allocation serves every + decode step. msa_proxy_max_score_view slices the per-step shape out + of it. """ buffers = self.cuda_graph_buffers self.msa_max_score = self.get_empty( buffers, - (num_index_heads * max_k_tiles * max_batch,), + (num_index_heads * max_k_tiles * max_tokens,), cache_name="msa_max_score", dtype=torch.float32, capture_graph=capture_graph, ) self.msa_n_valid_blocks = self.get_empty( buffers, - (max_batch,), + (max_tokens,), cache_name="msa_n_valid_blocks", dtype=torch.int32, capture_graph=capture_graph, @@ -430,14 +448,15 @@ def _ensure_msa_decode_scratch_buffers( required_max_k_tiles: int, ) -> None: """Ensure proxy scratch buffers exist and cover the current plan.""" - required_numel = num_index_heads * required_max_k_tiles * max_batch + max_tokens = max(int(max_batch), self._msa_max_decode_tokens()) + required_numel = num_index_heads * required_max_k_tiles * max_tokens if self.msa_max_score is not None: if self.msa_max_score.numel() < required_numel: raise ValueError( f"msa_max_score backing store ({self.msa_max_score.numel()} " f"elements) is smaller than the decode plan needs " f"({required_numel} = {num_index_heads} heads * " - f"{required_max_k_tiles} k-tiles * {max_batch} batch)." + f"{required_max_k_tiles} k-tiles * {max_tokens} tokens)." ) return @@ -459,7 +478,7 @@ def _ensure_msa_decode_scratch_buffers( ) self._alloc_msa_proxy_scratch( num_index_heads=num_index_heads, - max_batch=max_batch, + max_tokens=max_tokens, max_k_tiles=max_k_tiles, capture_graph=capture_graph, ) @@ -614,7 +633,10 @@ def _build_step_plans(self) -> None: n_valid = per_token_valid_blocks( qo_lens_cpu, kv_lens_cpu, qo_offset_cpu, causal=True, block_size=page_size ) - self.msa_n_valid_blocks[:batch].copy_(n_valid.to(torch.int32), non_blocking=True) + # Per-TOKEN entries: under speculative multi-token verify the decode + # step carries qo_len > 1 query tokens per request. + total_q = int(n_valid.shape[0]) + self.msa_n_valid_blocks[:total_q].copy_(n_valid.to(torch.int32), non_blocking=True) def _build_msa_fields(self) -> None: """Populate the MSA cache-write buffers for this step. @@ -640,13 +662,11 @@ def _build_msa_fields(self) -> None: cache_device = _cache_device(self) page_size = int(kv_cache_manager.tokens_per_block) - is_prefill = int(self.num_contexts or 0) > 0 - if not is_prefill and int(qo_lens_cpu.max().item()) > 1: - raise NotImplementedError( - "MiniMax-M3 MSA attention does not support speculative decoding " - "(multiple query tokens per decode step). Disable speculative " - "decoding or use the non-MSA MiniMax-M3 backend." - ) + # Multi-token generation rows (one-model Eagle3 spec verify emits + # 1 + draft_len query tokens per gen row) are expressed naturally: + # qo_lens carries the per-request query counts and every consumer + # below (fmha_sm100 plans, slot mapping, per-token valid blocks) is + # varlen. Decode batches therefore stay decode-shaped under spec. # Built in prepare() (outside capture), so these transients are # fine: forwards read only the persistent buffers filled below. diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py index 81b12cd5bce6..0abb1af37583 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py @@ -94,6 +94,13 @@ class MiniMaxM3TritonSparseAttentionMetadata: prefix_lens: Optional[torch.Tensor] = None cu_seqlens_q: Optional[torch.Tensor] = None extend_seq_lens_cpu: Optional[List[int]] = None + # Query tokens per request on decode-shaped metadata (``is_prefill=False``). + # 1 for ordinary decode; 1 + draft_len under one-model Eagle3 spec verify + # when a pure-generation batch is kept decode-shaped. The reference + # backend routes multi-token gen batches through the extend path, so this + # stays 1 there; the dense SDPA decode branch consumes it for the causal + # ladder mask. + decode_qo_len: int = 1 q_batch_row: Optional[torch.Tensor] = None q_positions: Optional[torch.Tensor] = None max_seqlen_q: int = field(default=1) @@ -866,7 +873,12 @@ def prepare(self) -> None: # specialization. (iter-131 regression: previously a wrong # predicate routed mixed batches into the decode branch and # crashed in index_copy_.) - is_extend = num_contexts > 0 + # Multi-token generation rows (one-model Eagle3 spec verify emits + # 1 + draft_len query tokens per gen row) route through the extend + # path: the prefill kernels handle them as prefix+window extends. + # The decode branch stays reserved for batches where every row + # appends exactly one token. + is_extend = num_contexts > 0 or int(seq_lens_cpu[:batch_size].max().item()) > 1 if is_extend: prefix_lens_list = [int(num_cached_per_seq[b]) for b in range(batch_size)] extend_seq_lens_cpu = [ diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index 461fbc981242..db339933bcbb 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -1145,8 +1145,7 @@ def _sdpa_dense_attention_core( # Build the per-query attention mask that masks out padded KV # positions beyond each sequence's true ``seq_lens`` and (for # prefill) preserves causality. ``q_positions`` from the prefill - # metadata names each Q token's K-side position; for decode - # there is one Q token per request at position ``seq_lens - 1``. + # metadata names each Q token's K-side position. # The metadata tensors are produced by # :meth:`MiniMaxM3AttentionMetadata.prepare` on the cache # device, so ``.to(dtype=torch.long)`` is a same-device dtype @@ -1195,14 +1194,22 @@ def _sdpa_dense_attention_core( ) # [1, H, q, d] output_view[start:end].copy_(out_b.squeeze(0).transpose(0, 1)) else: - # Decode: one Q token per request at position seq_lens - 1. - # Every input tensor here is already on q.device (set up by - # prepare()), so SDPA captures cleanly. - valid = kv_positions < seq_lens_dev.unsqueeze(-1) # [batch, max_k] - q_b = q_view.unsqueeze(1).transpose(1, 2) # [batch, H, 1, d] + # Decode: qo_len query tokens per request; token t of request b + # attends seq_lens[b] - qo_len + t + 1 positions (the causal + # ladder; qo_len=1 is the classic one-token mask). Every input + # tensor here is already on q.device (set up by prepare()), so + # SDPA captures cleanly. + qo_len = int(m3_meta.decode_qo_len) + ladder = torch.arange(1 - qo_len, 1, device=q.device, dtype=torch.long) + # eff[b, t] = attendable position count for token t of row b. + eff = seq_lens_dev.unsqueeze(-1) + ladder # [batch, qo] + valid = kv_positions.unsqueeze(1) < eff.unsqueeze(-1) # [batch, qo, max_k] + q_b = q_view.view(batch, qo_len, self.num_heads, self.head_dim).transpose( + 1, 2 + ) # [batch, H, qo, d] k_b = k_padded.transpose(1, 2) # [batch, H, k, d] v_b = v_padded.transpose(1, 2) # [batch, H, k, d] - mask_b = valid.unsqueeze(1).unsqueeze(1) # [batch, 1, 1, k] + mask_b = valid.unsqueeze(1) # [batch, 1, qo, k] with sdpa_kernel(_DENSE_SDPA_BACKENDS): out_b = torch.nn.functional.scaled_dot_product_attention( q_b.to(q.dtype), @@ -1211,19 +1218,11 @@ def _sdpa_dense_attention_core( attn_mask=mask_b, dropout_p=0.0, is_causal=False, - ) # [batch, H, 1, d] - # Drop the singleton Q-length axis and write the resulting - # ``[batch, num_heads, head_dim]`` tensor into the final buffer. - # The prior ``.transpose(1, 2).reshape(batch, H, d)`` pattern - # was wrong: with ``H != head_dim`` (M3 TP=8 has H=8, d=128) - # the non-contiguous transpose forces ``reshape`` to copy the - # data in C-order under its current ``[batch, d, H]`` shape, - # then reinterpret as ``[batch, H, d]`` — which scrambles - # ``(head, head_dim)`` ordering and feeds permuted activations - # into ``o_proj``. Prefill is unaffected because its - # ``transpose(0, 1)`` runs between q-len and num_heads axes - # which the per-batch loop already laid out correctly. - output.view(batch, self.num_heads, self.head_dim).copy_(out_b.squeeze(2)) + ) # [batch, H, qo, d] + # Copy through a token-major [batch, qo, H, dh] view rather + # than transpose(1, 2).reshape, which (with H != head_dim) + # copies in C-order and scrambles (head, head_dim) into o_proj. + output.view(batch, qo_len, self.num_heads, self.head_dim).copy_(out_b.transpose(1, 2)) return output @@ -1705,7 +1704,6 @@ def forward( hidden_states = inputs_embeds residual = None -<<<<<<< HEAD for layer_idx, decoder_layer in enumerate(self.layers): # Per-layer NVTX range (layer0, layer1, ...). Emitted only when # TLLM_NVTX_DEBUG=1 (or TLLM_LLMAPI_ENABLE_NVTX=1) is set. @@ -1715,17 +1713,8 @@ def forward( hidden_states=hidden_states, attn_metadata=attn_metadata, residual=residual, + spec_metadata=spec_metadata, ) -======= - for decoder_layer in self.layers: - hidden_states, residual = decoder_layer( - position_ids=position_ids, - hidden_states=hidden_states, - attn_metadata=attn_metadata, - residual=residual, - spec_metadata=spec_metadata, - ) ->>>>>>> bad119b1ce ([TRTLLM-14093][feat] Enable one-model Eagle3 speculative decoding for MiniMax-M3) hidden_states, _ = self.norm(hidden_states, residual) return hidden_states diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 89aa328e00ae..89794a8e0cee 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -676,21 +676,6 @@ def drafting_loop_wrapper(model): max_num_tokens = model_engine.max_num_tokens sparse_attention_config = model_engine.sparse_attention_config - # The MSA kernel path packs one query token per generation row and its - # prefill plan stages CPU values that go stale under overlap kv-len - # correction, so it cannot verify draft tokens (two-model spec also - # emits multi-token target rows) — reject at creation instead of at the - # first verify step inside prepare(). - if (sparse_attention_config is not None - and sparse_attention_config.algorithm == "minimax_m3" - and getattr(sparse_attention_config, "sparse_use_msa", False) - and spec_config is not None): - raise NotImplementedError( - "Speculative decoding is not supported with the MiniMax-M3 MSA " - "kernel path (sparse_use_msa=True): MSA packs one query token " - "per generation row. Set sparse_use_msa=False to use the " - "reference sparse backend with speculative decoding.") - if (sparse_attention_config is not None and sparse_attention_config.algorithm == "minimax_m3" and spec_config is not None diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index aac089c1c872..84c9309601f9 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -7666,11 +7666,16 @@ def test_nvfp4(self, use_msa): @pytest.mark.skip_less_device(4) @pytest.mark.skip_less_device_memory(140000) + @parametrize_with_ids("use_msa", [True]) @parametrize_with_ids("overlap_scheduler", [True]) @parametrize_with_ids("attention_dp", [False, True]) @parametrize_with_ids("tp_size,ep_size", [(4, 4)]) def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, - overlap_scheduler): + overlap_scheduler, use_msa): + if use_msa: + pytest.importorskip( + "fmha_sm100", + reason="MSA kernels (fmha_sm100) not installed") model_name = "nvidia/MiniMax-M3-NVFP4" model_path = f"{llm_models_root()}/MiniMax-M3-NVFP4" max_draft_len = 3 @@ -7678,13 +7683,16 @@ def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, max_draft_len=max_draft_len, speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3", ) + # The MSA kernels require page_size == sparse_block_size (128). kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.6, - enable_block_reuse=False) + enable_block_reuse=False, + tokens_per_block=128 if use_msa else 32) with LLM(model_path, tensor_parallel_size=tp_size, moe_expert_parallel_size=ep_size, kv_cache_config=kv_cache_config, - sparse_attention_config=MiniMaxM3SparseAttentionConfig(), + sparse_attention_config=MiniMaxM3SparseAttentionConfig( + sparse_use_msa=use_msa), moe_config=MoeConfig(backend="CUTLASS"), max_seq_len=4096, speculative_config=spec_config, diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index 5d8e877b892d..9dd42aca1b81 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -689,8 +689,8 @@ accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8[use_msa=True] TIMEOU accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8_piecewise_cuda_graph[tp_size=8-ep_size=8] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=False] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=True] TIMEOUT (180) -accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=False-overlap_scheduler=True] TIMEOUT (180) -accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=True-overlap_scheduler=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=False-overlap_scheduler=True-use_msa=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=True-overlap_scheduler=True-use_msa=True] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_auto_dtype accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_fp8 accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] 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 2e44c27e9bd3..624ba4b19ff9 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 @@ -223,3 +223,49 @@ def test_msa_proxy_max_score_strided_index_k_matches_packed(): assert not index_k_strided.is_contiguous() assert index_k_strided.stride(0) == coalescing_scale * page_size * head_dim assert torch.equal(strided_scores, packed_scores) + +def test_msa_scratch_sizing_covers_spec_verify_tokens(): + """Under one-model Eagle3 spec verify a decode step carries + 1 + draft_len query tokens per request, so the proxy scratch must be + sized by the worst-case decode TOKEN count, not the batch size. + """ + metadata_cls = MiniMaxM3MsaSparseAttention.Metadata + metadata = metadata_cls.__new__(metadata_cls) + metadata.kv_cache_manager = None + # 2 sequences, 4 tokens each (draft_len=3): 8 decode tokens per step. + metadata.max_num_sequences = 2 + metadata.max_num_tokens = 8 + # Store sized for batch-only sizing (4 heads * 16 k-tiles * 2), which is + # too small once tokens are accounted for (4 * 16 * 8). + metadata.msa_max_score = torch.zeros(4 * 16 * 2) + + with pytest.raises(ValueError, match=r"msa_max_score backing store"): + metadata._ensure_msa_decode_scratch_buffers( + num_index_heads=4, + max_batch=2, + capture_graph=False, + required_max_k_tiles=16, + ) + + +def test_per_token_valid_blocks_multi_token_decode(): + """Spec-verify decode rows expose one entry per query TOKEN, walking the + causal ladder within the verify window.""" + from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import ( + per_token_valid_blocks, ) + + # One request verifying 4 tokens against kv_len 10 (offset 6): token t + # attends 7 + t positions; with 2-token blocks that is ceil((7+t)/2). + qo = torch.tensor([4], dtype=torch.int32) + kv = torch.tensor([10], dtype=torch.int32) + off = torch.tensor([6], dtype=torch.int32) + n_valid = per_token_valid_blocks(qo, kv, off, causal=True, block_size=2) + assert n_valid.tolist() == [4, 4, 5, 5] + + # Mixed batch: an ordinary decode row (qo=1) alongside a verify row. + qo = torch.tensor([1, 3], dtype=torch.int32) + kv = torch.tensor([9, 6], dtype=torch.int32) + off = kv - qo + n_valid = per_token_valid_blocks(qo, kv, off, causal=True, block_size=4) + # Row 0: 9 positions -> 3 blocks. Row 1 tokens attend 4, 5, 6 -> 1, 2, 2. + assert n_valid.tolist() == [3, 1, 2, 2] From 667fd98f098e6c718b5433ea0a61e66ff238ad60 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Fri, 10 Jul 2026 19:07:32 +0000 Subject: [PATCH 4/5] [TRTLLM-14093][feat] Enable CUDA graphs for MiniMax-M3 MSA + Eagle3 The MSA decode driver was built capture-safe (device-only replan, prepare-time buffer allocation, stable data_ptrs), so multi-token verify no longer routes through the eager extend path on the MSA backend. Allow cuda_graph_config with sparse_use_msa=True; the reference path keeps raising. - speculative/eagle3.py: save/restore kv_lens_cuda across CUDA-graph warmup iterations. The one-model worker saved _seq_lens and the spec-decoding tensors, but not kv_lens_cuda, which the draft loop mutates in place - so the runner's pre-capture warmup iterations ran with drifted kv lens. Same pattern dflash/pard use; active only during graph warmup, not capture. - py_executor_creator.py: the graphs+spec rejection now applies only to the reference path (its verify routes through the eager extend path). - Capture hardening: the dense layers 0-2 decode branch now expands GQA K/V per KV head instead of all heads at once (bitwise-identical math; the expansion is captured into the graph pool, and under attention DP - unsharded heads - the whole-tensor transient exceeds the pool budget at large graph buckets). They also baked the prepare-time host upper bound max_seqlen_k into the captured gather/mask width; replays whose kv_len outgrew the capture-time value would silently truncate attention. Under graphs, bake min(page-table capacity, engine max_seq_len) instead (raw capacity alone OOMs during KV estimation; the seq_lens mask already invalidates positions past each row's true length). resolve_decode_state now raises if it would BUILD a decode state (JIT + allocation) while the stream is capturing. - test_nvfp4_eagle3: cuda_graph single-choice parametrize; the gated variants run the endgame config (MSA + Eagle3 + overlap + graphs). Support matrix updated (EAGLE-3 Linear: Yes). Probe (4xB200, TP4/EP4, batch-1, overlap, greedy): 25.98 tok/s with graphs vs 13.70 eager (1.90x), AR/AL in band; accuracy gate under graphs (adp=False): MMLU 84.62 / GSM8K 90.86, AR 0.719 / AL 3.156 - identical to eager; spec-dec graph capture confirmed on all ranks (draft_len=3 buckets); negative control: reference path + graphs + spec still raises at creation. Signed-off-by: Zheyu Fu (cherry picked from commit c8b09eb07e536b54d0dbd4995b05f002e6dceb4b) Signed-off-by: Zheyu Fu --- docs/source/models/supported-models.md | 4 +- .../_torch/models/modeling_minimaxm3.py | 53 ++++++++++++++----- .../_torch/pyexecutor/py_executor_creator.py | 13 +++-- tensorrt_llm/_torch/speculative/eagle3.py | 17 ++++-- .../defs/accuracy/test_llm_api_pytorch.py | 41 ++++++++------ .../test_lists/qa/llm_function_core.txt | 4 +- 6 files changed, 89 insertions(+), 43 deletions(-) diff --git a/docs/source/models/supported-models.md b/docs/source/models/supported-models.md index 46f90a9df31c..7e05e60017fa 100644 --- a/docs/source/models/supported-models.md +++ b/docs/source/models/supported-models.md @@ -78,7 +78,7 @@ Note: Support for other models may vary. Features marked "N/A" are not applicabl | `Gemma4ForConditionalGeneration` | Untested | Yes | Untested | No | Yes | No | No | No | No | Yes | Untested | No | Yes | Untested | Untested | | `Gemma4UnifiedForConditionalGeneration` | Untested | Untested | Untested | No | Yes | No | No | No | No | Yes | Untested | No | Yes | Untested | Untested | | `Step3p7ForConditionalGeneration`| Yes | Yes | Yes | Untested | Untested | Yes | No | No | No | Yes | Untested | Untested | Yes | Untested | Untested | -| `MiniMaxM3SparseForConditionalGeneration` [^12] | Yes | Yes | Yes | Untested | Untested | No | No | No | No | Yes | Untested | No | N/A | Untested | Untested | +| `MiniMaxM3SparseForConditionalGeneration` [^12] | Yes | Yes | Yes | Untested | Untested | No | Yes | No | No | Yes | Untested | No | N/A | Untested | Untested | [^1]: Chunked Prefill for MLA can only be enabled on SM100/SM103. [^2]: KV cache reuse for MLA can only be enabled on SM90/SM100/SM103 and in BF16/FP8 KV cache dtype. @@ -90,7 +90,7 @@ Note: Support for other models may vary. Features marked "N/A" are not applicabl [^9]: Audio modality only supported on E2B/E4B variants. [^10]: Audio requires a checkpoint with a `sound_config` and is supported only on the full (non-disaggregated) model path, not the EPD disaggregated path. [^11]: DeepSeek-V4 is only supported on Blackwell GPUs (`SM100+`). See the [DeepSeek-V4 example README](../../../examples/models/core/deepseek_v4/README.md) for setup and parallelism. -[^12]: Supports text, image, and video inputs over the block-sparse attention path. The published MXFP8 checkpoint is dequantized on load so the runtime sees an effectively BF16 model. The text decoder is also usable standalone (text-only) via the `MiniMaxM3SparseForCausalLM` architecture. KV cache reuse and MTP are not supported on the sparse-attention path in this release. +[^12]: Supports text, image, and video inputs over the block-sparse attention path. The published MXFP8 checkpoint is dequantized on load so the runtime sees an effectively BF16 model. The text decoder is also usable standalone (text-only) via the `MiniMaxM3SparseForCausalLM` architecture. KV cache reuse and MTP are not supported on the sparse-attention path in this release. One-model linear EAGLE-3 is supported; combining it with CUDA graphs requires the MSA kernels (`sparse_use_msa=True`, SM100). [^13]: The Cosmos 3 family also supports visual generation through the VisualGen API. See [Visual Generation Models](#visual-generation-models). # Multimodal Feature Support Matrix (PyTorch Backend) diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index db339933bcbb..db94b23fcc61 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -1111,7 +1111,22 @@ def _sdpa_dense_attention_core( # 7. Gather padded K/V for every batch row and run dense GQA. batch = int(m3_meta.slot_ids.shape[0]) - max_k = int(m3_meta.max_seqlen_k) + # Under CUDA-graph capture the gather/mask width is baked into the + # graph, while max_seqlen_k is a prepare-time host upper bound that + # later replays can outgrow (silent attention truncation). Bake the + # static bound instead: no row's kv can exceed the engine + # max_seq_len (manager-derived, includes the spec-dec margin), and + # the page-table width caps it when smaller. The raw page-table + # width alone is NOT usable — the KV-estimation pass inflates it + # far past max_seq_len and the [batch, max_k, heads] gather would + # OOM. The seq_lens mask below invalidates positions past each + # row's true length. + if getattr(attn_metadata, "is_cuda_graph", False): + capacity = int(m3_meta.req_to_token.shape[1]) + engine_bound = int(getattr(attn_metadata, "max_seq_len", None) or capacity) + max_k = min(capacity, engine_bound) + else: + max_k = int(m3_meta.max_seqlen_k) if max_k <= 0: max_k = 1 # ``_gather_paged_batched`` decomposes the flat slot id into @@ -1138,9 +1153,6 @@ def _sdpa_dense_attention_core( f"by num_key_value_heads ({self.num_key_value_heads})" ) group = self.num_heads // max(self.num_key_value_heads, 1) - if group > 1: - k_padded = k_padded.repeat_interleave(group, dim=2) - v_padded = v_padded.repeat_interleave(group, dim=2) # Build the per-query attention mask that masks out padded KV # positions beyond each sequence's true ``seq_lens`` and (for @@ -1154,6 +1166,9 @@ def _sdpa_dense_attention_core( kv_positions = torch.arange(max_k, device=q.device).unsqueeze(0) # [1, max_k] if m3_meta.is_prefill: + if group > 1: + k_padded = k_padded.repeat_interleave(group, dim=2) + v_padded = v_padded.repeat_interleave(group, dim=2) # Prefill: build [total_q, max_k] mask using q_positions / q_batch_row. # Prefill never runs inside the CUDA-graph capture window # (capture is decode-only), so the per-batch Python loop and @@ -1207,18 +1222,28 @@ def _sdpa_dense_attention_core( q_b = q_view.view(batch, qo_len, self.num_heads, self.head_dim).transpose( 1, 2 ) # [batch, H, qo, d] - k_b = k_padded.transpose(1, 2) # [batch, H, k, d] - v_b = v_padded.transpose(1, 2) # [batch, H, k, d] mask_b = valid.unsqueeze(1) # [batch, 1, qo, k] + # Expand K/V per KV head rather than all heads at once: this + # branch is CUDA-graph captured, so the expansion lives in the + # graph pool, and the full-head copy is O(batch * max_k * + # num_heads) - under attention DP (unsharded heads) that one + # transient exceeds the pool budget at large graph buckets. + # Per-head chunks are freed between iterations; with TP-sharded + # KV heads (1 per rank) the loop is a single iteration. + out_b = q.new_empty(batch, self.num_heads, qo_len, self.head_dim) with sdpa_kernel(_DENSE_SDPA_BACKENDS): - out_b = torch.nn.functional.scaled_dot_product_attention( - q_b.to(q.dtype), - k_b.to(q.dtype), - v_b.to(q.dtype), - attn_mask=mask_b, - dropout_p=0.0, - is_causal=False, - ) # [batch, H, qo, d] + for h in range(max(self.num_key_value_heads, 1)): + qh = slice(h * group, (h + 1) * group) + k_h = k_padded[:, :, h : h + 1].repeat_interleave(group, dim=2) + v_h = v_padded[:, :, h : h + 1].repeat_interleave(group, dim=2) + out_b[:, qh] = torch.nn.functional.scaled_dot_product_attention( + q_b[:, qh].to(q.dtype), + k_h.transpose(1, 2).to(q.dtype), + v_h.transpose(1, 2).to(q.dtype), + attn_mask=mask_b, + dropout_p=0.0, + is_causal=False, + ) # [batch, group, qo, d] # Copy through a token-major [batch, qo, H, dh] view rather # than transpose(1, 2).reshape, which (with H != head_dim) # copies in C-order and scrambles (head, head_dim) into o_proj. diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 89794a8e0cee..99ab707beef7 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -687,12 +687,15 @@ def drafting_loop_wrapper(model): "attention: the M3 sparse kernels implement linear-chain " "verification only. Remove eagle_choices / use_dynamic_tree " "from the speculative config.") - if llm_args.cuda_graph_config is not None: + if (llm_args.cuda_graph_config is not None + and not sparse_attention_config.sparse_use_msa): raise NotImplementedError( - "CUDA graphs are not supported with MiniMax-M3 sparse " - "attention and speculative decoding: multi-token verify " - "routes through the M3 extend path, which is not " - "capture-safe yet. Set cuda_graph_config to null.") + "CUDA graphs are not supported with MiniMax-M3 reference " + "sparse attention and speculative decoding: multi-token " + "verify routes through the M3 extend path, which is not " + "capture-safe. Set sparse_use_msa=True (SM100) to run " + "verify through the capture-safe MSA decode driver, or set " + "cuda_graph_config to null.") if not should_use_separate_draft_kv_cache(spec_config): raise NotImplementedError( "One-model speculative decoding with MiniMax-M3 sparse " diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index d9b5be1ecc1e..025c5fd6e024 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -628,8 +628,19 @@ def __init__(self, def max_draft_len(self) -> int: return self.spec_config.max_draft_len - def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): - attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda") + def _prepare_attn_metadata_for_spec_dec(self, attn_metadata, spec_metadata): + # During CUDA-graph warmup (graph metadata, not capturing), also + # save/restore kv_lens_cuda: the drafting loop mutates it in place + # and the runner warms up twice before capture, so without the + # restore the second warmup and the capture run with drifted kv + # lens (same pattern as DFlash/PARD). During capture itself the + # mutation must be captured, so kv_lens_cuda is not saved there. + is_capturing = torch.cuda.is_current_stream_capturing() + if spec_metadata.is_cuda_graph and not is_capturing: + attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda", + "kv_lens_cuda") + else: + attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda") batch_size = attn_metadata.num_seqs # Save spec-dec params that the drafting loop will overwrite. @@ -737,7 +748,7 @@ def forward(self, ) # Save the old attn_metadata and spec_metadata - self._prepare_attn_metadata_for_spec_dec(attn_metadata) + self._prepare_attn_metadata_for_spec_dec(attn_metadata, spec_metadata) # Prepare inputs for the 1st draft model forward position_ids = position_ids.squeeze(0) diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 84c9309601f9..35ec03b54a6f 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -7666,16 +7666,16 @@ def test_nvfp4(self, use_msa): @pytest.mark.skip_less_device(4) @pytest.mark.skip_less_device_memory(140000) + @parametrize_with_ids("cuda_graph", [True]) @parametrize_with_ids("use_msa", [True]) @parametrize_with_ids("overlap_scheduler", [True]) @parametrize_with_ids("attention_dp", [False, True]) @parametrize_with_ids("tp_size,ep_size", [(4, 4)]) def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, - overlap_scheduler, use_msa): + overlap_scheduler, use_msa, cuda_graph): if use_msa: - pytest.importorskip( - "fmha_sm100", - reason="MSA kernels (fmha_sm100) not installed") + pytest.importorskip("fmha_sm100", + reason="MSA kernels (fmha_sm100) not installed") model_name = "nvidia/MiniMax-M3-NVFP4" model_path = f"{llm_models_root()}/MiniMax-M3-NVFP4" max_draft_len = 3 @@ -7687,19 +7687,26 @@ def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.6, enable_block_reuse=False, tokens_per_block=128 if use_msa else 32) - with LLM(model_path, - tensor_parallel_size=tp_size, - moe_expert_parallel_size=ep_size, - kv_cache_config=kv_cache_config, - sparse_attention_config=MiniMaxM3SparseAttentionConfig( - sparse_use_msa=use_msa), - moe_config=MoeConfig(backend="CUTLASS"), - max_seq_len=4096, - speculative_config=spec_config, - cuda_graph_config=None, - disable_overlap_scheduler=not overlap_scheduler, - enable_attention_dp=attention_dp, - trust_remote_code=True) as llm: + with LLM( + model_path, + tensor_parallel_size=tp_size, + moe_expert_parallel_size=ep_size, + kv_cache_config=kv_cache_config, + sparse_attention_config=MiniMaxM3SparseAttentionConfig( + sparse_use_msa=use_msa), + moe_config=MoeConfig(backend="CUTLASS"), + max_seq_len=4096, + speculative_config=spec_config, + # Graphs + spec requires the MSA path: its verify batches + # are decode-shaped and capture-safe (the reference path + # rejects graphs+spec at creation). + cuda_graph_config=CudaGraphConfig( + enable_padding=True, + max_batch_size=64 if attention_dp else 128, + ) if cuda_graph else None, + disable_overlap_scheduler=not overlap_scheduler, + enable_attention_dp=attention_dp, + trust_remote_code=True) as llm: assert llm.args.quant_config.quant_algo == QuantAlgo.MIXED_PRECISION task = MMLU(model_name) task.evaluate(llm) diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index 9dd42aca1b81..b8de89d10aa7 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -689,8 +689,8 @@ accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8[use_msa=True] TIMEOU accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_mxfp8_piecewise_cuda_graph[tp_size=8-ep_size=8] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=False] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4[use_msa=True] TIMEOUT (180) -accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=False-overlap_scheduler=True-use_msa=True] TIMEOUT (180) -accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=True-overlap_scheduler=True-use_msa=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=False-overlap_scheduler=True-use_msa=True-cuda_graph=True] TIMEOUT (180) +accuracy/test_llm_api_pytorch.py::TestMiniMaxM3::test_nvfp4_eagle3[tp_size=4-ep_size=4-attention_dp=True-overlap_scheduler=True-use_msa=True-cuda_graph=True] TIMEOUT (180) accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_auto_dtype accuracy/test_llm_api_pytorch.py::TestMinistral8BInstruct::test_fp8 accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] From 499a3699ee424997c3d582db04a2b3a889308cf4 Mon Sep 17 00:00:00 2001 From: Zheyu Fu Date: Mon, 20 Jul 2026 06:41:07 +0000 Subject: [PATCH 5/5] [TRTLLM-14093][feat] Adapt Eagle3 support to the new side branch MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Squash of the follow-up work that makes the four ported Eagle3 commits functional and validated on feat/m3_with_msa: - Adapt to this branch's API: the creator guard and integration test used the removed sparse_use_msa field (implementation= now); MsaSparseGqaFmha.is_available crashed on the dense Eagle3 drafter (sparse_params=None). - Correct the MSA length sources: msa_kv_lens_cpu excludes num_extra_kv_tokens (nonzero only under speculation; inflated every slot/page/plan length and crashed spec warmup with an illegal memory access), and the graph-safe plan mirrors are sized per expanded token row (the planner splits qo_len > 1 requests per token). - Device-side overlap correction: on_update_kv_lens on the MSA metadata re-derives the KV-write slots, per-token valid-block counts, and the per-row kv_segment_lens/qo_offset plan-mirror entries from the corrected kv_lens_cuda — pure in-stream device ops, capture-safe, shrink-only. Eager mixed batches install corrected host lens and rebuild fields and eager plans. Without this the overlap scheduler (on by default) silently corrupts the KV layout on every partially rejected verify step (GSM8K 57.9 vs 90.4+ with the hook). - Draft KV page size: MiniMaxM3KVCacheManagerV2 declares draft_manager_tokens_per_block = 32, read generically at draft-manager creation — a scoped workaround for a base regression (the drafter's trtllm-gen generation kernels IMA at tokens_per_block=128; the MSA target requires 128, the dense draft layers have no page-size constraint). Remove once the kernel bug is fixed. - Tests: test_nvfp4_eagle3 runs the endgame config (MSA + Eagle3 + overlap + CUDA graphs, attention-DP on/off) with MMLU/GSM8K gates and a chat-format GSM8K acceptance check against the drafter head's published reference (floors 0.78 / 3.3; head card 0.839 / 3.518); QA-list rows updated; two unit tests for token-count scratch sizing and the per-token valid-block ladder. - Review feedback: creator-level M3 spec guards removed, ADP shared- layout opt-out documented on the manager class, complexity/comment pass (net -70 lines), pre-commit formatting. Validated on 4xB200 (TP4/EP4, NVFP4, draft_len=3, no env setup): adp=False MMLU 85.14 / GSM8K 90.94, chat acceptance 0.843 / 3.528; adp=True 84.89 / 90.37, chat acceptance 0.846 / 3.538; spec-off regression test_nvfp4 both variants pass; 42/42 unit tests; overlap throughput +6.2-6.4% at acceptance-matched batches 32/128. Signed-off-by: Zheyu Fu --- .../attention_backend/fmha/msa_sparse_gqa.py | 5 +- .../sparse/minimax_m3/cache_manager.py | 8 + .../sparse/minimax_m3/msa_backend.py | 176 +++++++++++++++--- .../sparse/minimax_m3/triton_metadata.py | 71 +++---- .../_torch/models/modeling_minimaxm3.py | 39 ++-- tensorrt_llm/_torch/pyexecutor/_util.py | 21 ++- .../_torch/pyexecutor/py_executor_creator.py | 32 +--- tensorrt_llm/_torch/speculative/eagle3.py | 10 +- .../defs/accuracy/test_llm_api_pytorch.py | 100 +++++----- .../sparse/test_minimax_m3_msa_backend.py | 4 +- 10 files changed, 280 insertions(+), 186 deletions(-) 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 3f8c6af4da98..5578f36e3ca8 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py @@ -198,8 +198,9 @@ def is_available(cls, attn: Optional["TrtllmAttention"] = None) -> bool: return False # Only the MiniMax-M3 MSA layer uses this library. Matching the lowered # sparse algorithm lets the base create_fmha_libs add it to that layer - # alone, so no create_fmha_libs override is needed. - return attn.sparse_params.algorithm == "minimax_m3" + # alone, so no create_fmha_libs override is needed. Dense layers (e.g. + # an Eagle3 draft model) have no sparse_params. + return attn.sparse_params is not None and attn.sparse_params.algorithm == "minimax_m3" def forward( self, 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 fda13446fe2b..fd0f72159f1f 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 @@ -156,6 +156,14 @@ class MiniMaxM3KVCacheManagerV2(KVCacheManagerV2): # (read by ``_should_create_separate_draft_kv_cache``). supports_shared_draft_layers = False + # Workaround: the Eagle3 drafter's trtllm-gen generation kernels hit an + # illegal memory access at tokens_per_block=128 (base regression), while + # the MSA target requires 128. The dense draft layers have no page-size + # constraint, so the separate draft manager runs at 32 (read by + # ``_create_one_model_draft_kv_cache_manager``). Remove once the kernel + # bug is fixed. + draft_manager_tokens_per_block = 32 + def __init__( self, *args, 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 bfa81832a181..7ebc4768bc0f 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 @@ -253,6 +253,11 @@ def __post_init__(self) -> None: super().__post_init__() params = self.sparse_metadata_params self._msa_params = params if isinstance(params, MiniMaxM3SparseMetadataParams) else None + # See on_update_kv_lens. + self._msa_live_batch = 0 + self._msa_live_total_q = 0 + self._msa_page_size = 0 + self._msa_corrected_kv_lens_cpu: Optional[torch.Tensor] = None self._create_msa_buffers() @property @@ -266,11 +271,23 @@ def msa_qo_lens_cpu(self) -> Optional[torch.Tensor]: @property def msa_kv_lens_cpu(self) -> Optional[torch.Tensor]: - """Per-request KV length, cached plus new tokens (host int32).""" + """Per-request KV length, cached plus new tokens (host int32). + + The base ``kv_lens`` includes ``num_extra_kv_tokens`` (speculative + draft-loop slots consumed by the C++ kernels); the MSA plans, ladder + slots and page counts need the true attended length, so it is + excluded here. + """ + if self._msa_corrected_kv_lens_cpu is not None: + return self._msa_corrected_kv_lens_cpu kv_lens = getattr(self, "kv_lens", None) if self.seq_lens is None or kv_lens is None: return None out = kv_lens[: self.num_seqs] + params = self.kv_cache_params + extra = params.num_extra_kv_tokens if params is not None else 0 + if extra: + out = out - extra return out if out.dtype == torch.int32 else out.to(torch.int32) @property @@ -370,6 +387,37 @@ def _create_msa_buffers(self) -> None: dtype=torch.int32, capture_graph=capture_graph, ) + # Staging for on_update_kv_lens: re-derives slots/bounds on device + # from the corrected kv_lens_cuda, sync-free. + tokens_per_block = int(kv_cache_manager.tokens_per_block) + self.msa_req_to_token = self.get_empty( + buffers, + (max_num_sequences, max_blocks_per_seq * tokens_per_block), + cache_name="msa_req_to_token", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.msa_q_batch_row = self.get_empty( + buffers, + (max_num_tokens,), + cache_name="msa_q_batch_row", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.msa_q_intra = self.get_empty( + buffers, + (max_num_tokens,), + cache_name="msa_q_intra", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.msa_qo_lens_dev = self.get_empty( + buffers, + (max_num_sequences,), + cache_name="msa_qo_lens_dev", + dtype=torch.int32, + capture_graph=capture_graph, + ) # The proxy scratch needs the fmha_sm100 plan geometry. This metadata # exists only for the MSA backend, whose selection already required the # kernels, so a failed import here is a hard error rather than a reason @@ -392,19 +440,16 @@ def _create_msa_buffers(self) -> None: self._msa_buffers_ready = True def _msa_max_decode_tokens(self) -> int: - """Worst-case decode-step query-token count for scratch sizing. - - Non-spec decode has one query token per sequence; one-model Eagle3 - spec verify has 1 + draft_len per gen row. max_num_tokens bounds the - step's total token count in both cases and keeps the scratch a few - MB, so it is used directly rather than plumbing the draft length. - Falls back to 0 on partially built metadata (structural tests); - callers floor the result with their own batch/token counts. + """Worst-case decode-step query tokens (spec verify emits 1 + draft_len + per row), bounded by max_num_tokens. getattr fallbacks cover metadata + built via ``__new__`` in structural tests. """ max_seqs = int(getattr(self, "max_num_sequences", 0) or 0) max_toks = int(getattr(self, "max_num_tokens", 0) or 0) if max_toks <= 0: return max_seqs + # 16384 keeps total_q * num_qo_heads under the fmha_sm100 planner cap + # (65536, fmha_sm100/api.py) at 4 sharded index heads. return max(max_seqs, min(max_toks, 16384)) def _alloc_msa_proxy_scratch( @@ -485,9 +530,81 @@ def _ensure_msa_decode_scratch_buffers( def prepare(self) -> None: super().prepare() + self._msa_corrected_kv_lens_cpu = None self._build_msa_fields() self._build_step_plans() + def on_update_kv_lens(self) -> None: + """Re-derive length-dependent MSA state from the corrected kv_lens_cuda. + + The overlap scheduler corrects kv_lens_cuda on device after prepare() + staged optimistic (full-acceptance) lens. Decode steps are patched + with pure device ops (capture-safe); the correction only shrinks + lengths, so the host-baked plan worklists and the page-table layout + stay valid. Eager mixed batches install corrected host lens and + rebuild (small D2H sync). Idempotent. + """ + super().on_update_kv_lens() + if not self._msa_fields_ready: + return + batch = self._msa_live_batch + total_q = self._msa_live_total_q + if batch <= 0 or total_q <= 0: + return + if self.msa_decode_proxy_plan is None: + # Prefill/mixed step: the eager plans and the page-table layout + # were built from the optimistic host mirrors, so install the + # corrected lens and rebuild both. Mixed steps are never captured. + if torch.cuda.is_current_stream_capturing(): + return + self._msa_corrected_kv_lens_cpu = self.kv_lens_cuda[:batch].to("cpu", torch.int32) + self._build_msa_fields() + self._build_step_plans() + return + + kv_true = self.kv_lens_cuda[:batch] + qbr = self.msa_q_batch_row[:total_q].to(torch.long) + qo_dev = self.msa_qo_lens_dev[:batch] + kv_true_tok = kv_true[qbr] + pos = kv_true_tok - qo_dev[qbr] + self.msa_q_intra[:total_q] + + # KV/idx-K write slots: slot[j] = req_to_token[request[j], pos[j]]. + width = int(self.msa_req_to_token.shape[1]) + idx = pos.to(torch.long).clamp(min=0, max=width - 1) + slots = self.msa_req_to_token.reshape(-1).index_select(0, qbr * width + idx) + self.msa_out_cache_loc[:total_q].copy_(slots) + + # Per-token valid-block counts for top-k selection. clamp_min(1) + # keeps degenerate (graph-padding) rows from -inf-masking every + # block, which would NaN the fully-masked GQA row. + page = self._msa_page_size + n_valid = torch.div((pos + 1).clamp_min(1) + (page - 1), page, rounding_mode="floor") + self.msa_n_valid_blocks[:total_q].copy_(n_valid.to(torch.int32)) + + # Plan mirrors: proxy and dense keep one row per request; the sparse + # GQA plan is row-expanded per token with the ladder in the per-row + # offset. qo_offset must stay non-negative: negative values hit the + # kernel's packed-length sentinel fallback. + off_req = (kv_true - qo_dev).clamp_min(0) + for owner, expanded in ( + (self._msa_proxy_plan, False), + (self._msa_gqa_plan, True), + (self._msa_dense_plan, False), + ): + decode = owner.plan[3] + seg_lens = decode.get("kv_segment_lens") + qo_off = decode.get("qo_offset") + if expanded: + if seg_lens is not None: + seg_lens[:total_q].copy_(kv_true_tok.to(seg_lens.dtype)) + if qo_off is not None: + qo_off[:total_q].copy_(pos.clamp_min(0).to(qo_off.dtype)) + else: + if seg_lens is not None: + seg_lens[:batch].copy_(kv_true.to(seg_lens.dtype)) + if qo_off is not None: + qo_off[:batch].copy_(off_req.to(qo_off.dtype)) + def _build_step_plans(self) -> None: """Build the three layer-invariant fmha_sm100 plans once per step. @@ -528,7 +645,6 @@ def _build_step_plans(self) -> None: qo_offset_cpu = self.msa_qo_offset_cpu if qo_lens_cpu is None or kv_lens_cpu is None or qo_offset_cpu is None: return - batch = int(qo_lens_cpu.shape[0]) device = _cache_device(self) page_size = int(self.kv_cache_manager.tokens_per_block) capture_graph = self.is_cuda_graph @@ -599,27 +715,31 @@ def _build_step_plans(self) -> None: ) # Allocate the graph-safe plan owners once per metadata; later steps - # only refresh their contents below. + # only refresh their contents below. The plan worklists are sized per + # expanded row (the planner splits qo_len > 1 requests into per-token + # rows under speculative multi-token verify), so use the worst-case + # decode-step token count rather than the batch size. if self._msa_proxy_plan is None: + max_plan_rows = max(max_batch, self._msa_max_decode_tokens()) num_ctas = torch.cuda.get_device_properties(device).multi_processor_count self._msa_proxy_plan = _MsaGraphSafePlan( self, "msa_proxy_plan", - max_batch=max_batch, + max_batch=max_plan_rows, num_ctas=num_ctas, capture_graph=capture_graph, ) self._msa_gqa_plan = _MsaGraphSafePlan( self, "msa_gqa_plan", - max_batch=max_batch, + max_batch=max_plan_rows, num_ctas=num_ctas, capture_graph=capture_graph, ) self._msa_dense_plan = _MsaGraphSafePlan( self, "msa_dense_plan", - max_batch=max_batch, + max_batch=max_plan_rows, num_ctas=num_ctas, capture_graph=capture_graph, ) @@ -633,8 +753,7 @@ def _build_step_plans(self) -> None: n_valid = per_token_valid_blocks( qo_lens_cpu, kv_lens_cpu, qo_offset_cpu, causal=True, block_size=page_size ) - # Per-TOKEN entries: under speculative multi-token verify the decode - # step carries qo_len > 1 query tokens per request. + # One entry per query token (qo_len > 1 under spec verify). total_q = int(n_valid.shape[0]) self.msa_n_valid_blocks[:total_q].copy_(n_valid.to(torch.int32), non_blocking=True) @@ -662,12 +781,6 @@ def _build_msa_fields(self) -> None: cache_device = _cache_device(self) page_size = int(kv_cache_manager.tokens_per_block) - # Multi-token generation rows (one-model Eagle3 spec verify emits - # 1 + draft_len query tokens per gen row) are expressed naturally: - # qo_lens carries the per-request query counts and every consumer - # below (fmha_sm100 plans, slot mapping, per-token valid blocks) is - # varlen. Decode batches therefore stay decode-shaped under spec. - # Built in prepare() (outside capture), so these transients are # fine: forwards read only the persistent buffers filled below. # qo_offset is the prefix length, so one build covers prefill @@ -696,6 +809,25 @@ def _build_msa_fields(self) -> None: self.msa_out_cache_loc[:total_new_tokens].copy_(out_cache_loc, non_blocking=True) self.msa_kv_indices[:total_pages].copy_(kv_indices, non_blocking=True) + + # Staging for on_update_kv_lens. + step_width = int(req_to_token.shape[1]) + self.msa_req_to_token[:batch_size, :step_width].copy_(req_to_token, non_blocking=True) + qo_long = qo_lens_cpu.to(torch.long) + batch_row_cpu = torch.repeat_interleave( + torch.arange(batch_size, dtype=torch.int32), qo_long + ) + starts = torch.cumsum(qo_long, 0) - qo_long + intra_cpu = ( + torch.arange(total_new_tokens, dtype=torch.int64) + - torch.repeat_interleave(starts, qo_long) + ).to(torch.int32) + self.msa_q_batch_row[:total_new_tokens].copy_(batch_row_cpu) + self.msa_q_intra[:total_new_tokens].copy_(intra_cpu) + self.msa_qo_lens_dev[:batch_size].copy_(qo_lens_cpu) + self._msa_live_batch = batch_size + self._msa_live_total_q = total_new_tokens + self._msa_page_size = page_size self._msa_fields_ready = True def msa_idx_k_cache(self, layer_idx: int) -> torch.Tensor: diff --git a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py index 0abb1af37583..a57cf5eac2e1 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/minimax_m3/triton_metadata.py @@ -20,7 +20,6 @@ import torch -from ...interface import AttentionMetadata from ...trtllm import TrtllmAttentionMetadata from .common import build_paged_kv_slot_mapping @@ -94,12 +93,9 @@ class MiniMaxM3TritonSparseAttentionMetadata: prefix_lens: Optional[torch.Tensor] = None cu_seqlens_q: Optional[torch.Tensor] = None extend_seq_lens_cpu: Optional[List[int]] = None - # Query tokens per request on decode-shaped metadata (``is_prefill=False``). - # 1 for ordinary decode; 1 + draft_len under one-model Eagle3 spec verify - # when a pure-generation batch is kept decode-shaped. The reference - # backend routes multi-token gen batches through the extend path, so this - # stays 1 there; the dense SDPA decode branch consumes it for the causal - # ladder mask. + # Query tokens per request on decode-shaped metadata: 1 normally, + # 1 + draft_len under one-model Eagle3 verify. Consumed by the dense + # SDPA decode ladder mask. decode_qo_len: int = 1 q_batch_row: Optional[torch.Tensor] = None q_positions: Optional[torch.Tensor] = None @@ -175,10 +171,8 @@ def prepare(self) -> None: else: max_k = int(self.seq_lens_cpu[:batch_size].max().item()) # Optimistic overlap+spec lengths can overhang the page table at - # a page boundary (see derive_q_positions_and_cache_slots); the - # dense fallback consumes max_seqlen_k as the exact SDPA - # mask/gather width, so an unclamped overhang crashes SDPA - # (mask 385 vs K 384 on MMLU). + # a page boundary; SDPA consumes max_seqlen_k as the exact + # mask/gather width, so clamp to the table. self.max_seqlen_k = min(max_k, int(self.req_to_token.shape[1])) @@ -391,25 +385,18 @@ def derive_q_positions_and_cache_slots( ) -> Tuple[torch.Tensor, torch.Tensor]: """Per-Q-token K-side positions and KV slot ids, on device, sync-free. - ``q_positions[t] = prefix_lens[row(t)] + (t - cu_seqlens_q[row(t)])``; - ``out_cache_loc[t] = req_to_token[row(t), q_positions[t]]``. Shared by - the metadata builder and the ``on_update_kv_lens`` re-derivation so the - two cannot drift. + Shared by the metadata builder and the ``on_update_kv_lens`` + re-derivation so the two cannot drift. """ total_q = int(q_batch_row.shape[0]) qbr = q_batch_row.to(torch.long) tok = torch.arange(total_q, dtype=torch.int32, device=q_batch_row.device) q_positions = prefix_lens[qbr] + (tok - cu_seqlens_q[qbr]) - # Overlap+spec: the first derivation runs with optimistic - # (full-acceptance) prefix_lens, which overhang the last allocated page - # when a request's true length sits at a page boundary. The overhanging - # slots are placeholders (on_update_kv_lens re-derives them from the - # corrected kv_lens before any forward consumes them), but the gather - # must stay in bounds: unclamped, the last row trips the CUDA gather - # assert and inner rows silently read the next request's slots. Clamp - # the table index only (non-inplace: ``.to`` aliases int64 inputs) — - # q_positions keep the optimistic values, whose mask width prepare() - # bounds separately. + # Optimistic prefix_lens can overhang the last allocated page; the + # overhanging slots are placeholders (on_update_kv_lens re-derives them + # before any forward reads them) but the gather must stay in bounds. + # Clamp only the table index — non-inplace, since ``.to`` aliases int64 + # inputs. idx = q_positions.to(torch.long).clamp(min=0, max=req_to_token.shape[1] - 1) flat = qbr * req_to_token.shape[1] + idx return q_positions, req_to_token.reshape(-1).index_select(0, flat) @@ -428,8 +415,6 @@ def derive_decode_cache_slots(req_to_token: torch.Tensor, seq_lens: torch.Tensor return req_to_token.reshape(-1).index_select(0, flat) - - def build_runtime_metadata_from_kv_manager( *, kv_cache_manager, @@ -691,17 +676,10 @@ def build_runtime_metadata_from_kv_manager( class MiniMaxM3AttentionMetadata(TrtllmAttentionMetadata): """:class:`TrtllmAttentionMetadata` that pre-builds MiniMax-M3 metadata. - Subclassing :class:`TrtllmAttentionMetadata` (precedent: - ``DSAtrtllmAttentionMetadata``, and the MSA-path metadata) rather than - the plain :class:`AttentionMetadata` is required for one-model - speculative decoding (Eagle3): the draft layers live in the same engine - and run :class:`TrtllmAttention` against this shared per-step metadata, - so it must carry the TRTLLM surface (``kv_lens_cuda``, - ``host_request_types``, ``kv_cache_block_offsets``, - ``draft_kv_cache_block_offsets``, ``update_spec_dec_param``, ...). The - engine also gates spec-dec plumbing (draft KV cache swap, - overlap-scheduler ``kv_lens_cuda`` fixups) on - ``isinstance(..., TrtllmAttentionMetadata)``. + Subclasses :class:`TrtllmAttentionMetadata` (precedent: + ``DSAtrtllmAttentionMetadata``): one-model Eagle3 draft layers run + :class:`TrtllmAttention` against this shared per-step metadata, and the + engine gates spec-dec plumbing on the TRTLLM ``isinstance``. Overrides :meth:`prepare` so the M3-sparse :class:`MiniMaxM3TritonSparseAttentionMetadata` and the per-new-token @@ -873,11 +851,8 @@ def prepare(self) -> None: # specialization. (iter-131 regression: previously a wrong # predicate routed mixed batches into the decode branch and # crashed in index_copy_.) - # Multi-token generation rows (one-model Eagle3 spec verify emits - # 1 + draft_len query tokens per gen row) route through the extend - # path: the prefill kernels handle them as prefix+window extends. - # The decode branch stays reserved for batches where every row - # appends exactly one token. + # Multi-token gen rows (spec verify) also route through the extend + # path as prefix+window extends; decode stays one-token-per-row. is_extend = num_contexts > 0 or int(seq_lens_cpu[:batch_size].max().item()) > 1 if is_extend: prefix_lens_list = [int(num_cached_per_seq[b]) for b in range(batch_size)] @@ -951,13 +926,9 @@ def on_update_kv_lens(self) -> None: meta.q_positions[:total_q].copy_(q_positions) out_cache_loc[:total_q].copy_(cache_slots) else: - # Reached today only as an identity (the hook also fires - # pre-correction on ordinary decode steps); re-deriving - # keeps corrected 0-draft steps correct once dynamic - # draft lengths make them reachable. - out_cache_loc[:batch].copy_( - derive_decode_cache_slots(meta.req_to_token, kv_lens) - ) + # Identity today; keeps 0-draft corrections right if they + # become reachable. + out_cache_loc[:batch].copy_(derive_decode_cache_slots(meta.req_to_token, kv_lens)) __all__ = [ diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index db94b23fcc61..b31720dcea7e 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -71,12 +71,7 @@ is_torch_compiling, ) from .modeling_speculative import SpecDecOneEngineForCausalLM -from .modeling_utils import ( - DecoderModel, - ModelConfig, - filter_weights, - register_auto_model, -) +from .modeling_utils import DecoderModel, ModelConfig, filter_weights, register_auto_model # Dense layers use SDPA with non-contiguous Q/K/V and a bool attn_mask. # Limit backends to memory-efficient and math; cuDNN SDPA fails for this layout, @@ -1111,20 +1106,15 @@ def _sdpa_dense_attention_core( # 7. Gather padded K/V for every batch row and run dense GQA. batch = int(m3_meta.slot_ids.shape[0]) - # Under CUDA-graph capture the gather/mask width is baked into the - # graph, while max_seqlen_k is a prepare-time host upper bound that - # later replays can outgrow (silent attention truncation). Bake the - # static bound instead: no row's kv can exceed the engine - # max_seq_len (manager-derived, includes the spec-dec margin), and - # the page-table width caps it when smaller. The raw page-table - # width alone is NOT usable — the KV-estimation pass inflates it - # far past max_seq_len and the [batch, max_k, heads] gather would - # OOM. The seq_lens mask below invalidates positions past each - # row's true length. - if getattr(attn_metadata, "is_cuda_graph", False): + # Graph capture bakes the gather/mask width, and max_seqlen_k is a + # per-step host bound later replays can outgrow. Bake + # min(page-table width, engine max_seq_len): the raw table width + # alone is inflated far past max_seq_len by the KV-estimation pass + # and would OOM the [batch, max_k, heads] gather. The seq_lens mask + # below invalidates the slack. + if attn_metadata.is_cuda_graph: capacity = int(m3_meta.req_to_token.shape[1]) - engine_bound = int(getattr(attn_metadata, "max_seq_len", None) or capacity) - max_k = min(capacity, engine_bound) + max_k = min(capacity, int(attn_metadata.max_seq_len or capacity)) else: max_k = int(m3_meta.max_seqlen_k) if max_k <= 0: @@ -1223,13 +1213,10 @@ def _sdpa_dense_attention_core( 1, 2 ) # [batch, H, qo, d] mask_b = valid.unsqueeze(1) # [batch, 1, qo, k] - # Expand K/V per KV head rather than all heads at once: this - # branch is CUDA-graph captured, so the expansion lives in the - # graph pool, and the full-head copy is O(batch * max_k * - # num_heads) - under attention DP (unsharded heads) that one - # transient exceeds the pool budget at large graph buckets. - # Per-head chunks are freed between iterations; with TP-sharded - # KV heads (1 per rank) the loop is a single iteration. + # Expand K/V one KV head at a time: the all-heads transient is + # O(batch * max_k * num_heads) inside the CUDA-graph pool and + # exceeds the pool budget under attention DP (unsharded heads); + # with TP-sharded KV heads this is one iteration. out_b = q.new_empty(batch, self.num_heads, qo_len, self.head_dim) with sdpa_kernel(_DENSE_SDPA_BACKENDS): for h in range(max(self.num_key_value_heads, 1)): diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 1fee4a11808d..35fbfb71741f 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1048,8 +1048,12 @@ def _should_create_separate_draft_kv_cache(self) -> bool: if self._mapping.enable_attention_dp and getattr( self._kv_cache_manager_cls, 'supports_shared_draft_layers', True): - # Back-compat: attention DP keeps the shared-manager layout - # existing deployments were validated with. + # 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. logger.info("Attention DP: draft layers share the target KV " "cache manager.") return False @@ -1168,12 +1172,23 @@ 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=self._tokens_per_block, + tokens_per_block=draft_tpb, max_seq_len=self._max_seq_len, max_batch_size=self._max_batch_size, spec_config=self._speculative_config, diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 99ab707beef7..57aa45e606f6 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -35,8 +35,7 @@ from ..attention_backend.trtllm import TrtllmAttention from ..distributed import Distributed from ..speculative import (get_num_extra_kv_tokens, get_spec_drafter, - get_spec_resource_manager, - should_use_separate_draft_kv_cache) + get_spec_resource_manager) from ..virtual_memory import scope as virtual_memory_scope from ._util import (KvCacheCreator, _adjust_torch_mem_fraction, create_py_executor_instance, instantiate_sampler, is_mla, @@ -676,35 +675,6 @@ def drafting_loop_wrapper(model): max_num_tokens = model_engine.max_num_tokens sparse_attention_config = model_engine.sparse_attention_config - if (sparse_attention_config is not None - and sparse_attention_config.algorithm == "minimax_m3" - and spec_config is not None - and spec_config.spec_dec_mode.is_eagle3_one_model()): - if not spec_config.is_linear_tree: - raise NotImplementedError( - "Tree-based speculative decoding (eagle_choices / " - "use_dynamic_tree) is not supported with MiniMax-M3 sparse " - "attention: the M3 sparse kernels implement linear-chain " - "verification only. Remove eagle_choices / use_dynamic_tree " - "from the speculative config.") - if (llm_args.cuda_graph_config is not None - and not sparse_attention_config.sparse_use_msa): - raise NotImplementedError( - "CUDA graphs are not supported with MiniMax-M3 reference " - "sparse attention and speculative decoding: multi-token " - "verify routes through the M3 extend path, which is not " - "capture-safe. Set sparse_use_msa=True (SM100) to run " - "verify through the capture-safe MSA decode driver, or set " - "cuda_graph_config to null.") - if not should_use_separate_draft_kv_cache(spec_config): - raise NotImplementedError( - "One-model speculative decoding with MiniMax-M3 sparse " - "attention requires a separate draft KV cache manager, but " - "it is disabled for this configuration (e.g. disaggregated " - "serving disables it as a WAR for nvbug 5807902). Use " - "two-model speculative decoding (eagle3_one_model=False) " - "instead.") - # Set default value for cache_transceiver_config.max_tokens_in_buffer if cache_transceiver_config and cache_transceiver_config.max_tokens_in_buffer is None: cache_transceiver_config.max_tokens_in_buffer = net_max_seq_len diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 025c5fd6e024..d09bc91c3e8e 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -629,12 +629,10 @@ def max_draft_len(self) -> int: return self.spec_config.max_draft_len def _prepare_attn_metadata_for_spec_dec(self, attn_metadata, spec_metadata): - # During CUDA-graph warmup (graph metadata, not capturing), also - # save/restore kv_lens_cuda: the drafting loop mutates it in place - # and the runner warms up twice before capture, so without the - # restore the second warmup and the capture run with drifted kv - # lens (same pattern as DFlash/PARD). During capture itself the - # mutation must be captured, so kv_lens_cuda is not saved there. + # Graph warmup runs twice before capture while the draft loop + # mutates kv_lens_cuda in place — save/restore it during warmup + # (same as DFlash/PARD). During capture the mutation must be + # recorded, so skip the save there. is_capturing = torch.cuda.is_current_stream_capturing() if spec_metadata.is_cuda_graph and not is_capturing: attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda", diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 35ec03b54a6f..13a8ad792165 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -7668,14 +7668,16 @@ def test_nvfp4(self, use_msa): @pytest.mark.skip_less_device_memory(140000) @parametrize_with_ids("cuda_graph", [True]) @parametrize_with_ids("use_msa", [True]) - @parametrize_with_ids("overlap_scheduler", [True]) + @parametrize_with_ids("overlap_scheduler", [False, True]) @parametrize_with_ids("attention_dp", [False, True]) @parametrize_with_ids("tp_size,ep_size", [(4, 4)]) def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, overlap_scheduler, use_msa, cuda_graph): if use_msa: - pytest.importorskip("fmha_sm100", - reason="MSA kernels (fmha_sm100) not installed") + from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import \ + msa_package_available + if not msa_package_available(): + pytest.skip("MSA kernels (fmha_sm100) not available") model_name = "nvidia/MiniMax-M3-NVFP4" model_path = f"{llm_models_root()}/MiniMax-M3-NVFP4" max_draft_len = 3 @@ -7683,19 +7685,24 @@ def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, max_draft_len=max_draft_len, speculative_model=f"{llm_models_root()}/MiniMax-M3-EAGLE3", ) - # The MSA kernels require page_size == sparse_block_size (128). + # The runtime forces tokens_per_block per implementation (128 MSA / 32 + # reference). kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.6, - enable_block_reuse=False, - tokens_per_block=128 if use_msa else 32) + enable_block_reuse=False) with LLM( model_path, tensor_parallel_size=tp_size, moe_expert_parallel_size=ep_size, kv_cache_config=kv_cache_config, sparse_attention_config=MiniMaxM3SparseAttentionConfig( - sparse_use_msa=use_msa), + implementation="msa" if use_msa else "triton"), moe_config=MoeConfig(backend="CUTLASS"), max_seq_len=4096, + # The fmha_sm100 decode planner caps total_q x num_qo_heads at + # 65536; with 1 + draft_len = 4 verify tokens per row that + # bounds the batch at 1024 / TP-sharded heads (256 unsharded + # under attention DP). + max_batch_size=256 if attention_dp else 512, speculative_config=spec_config, # Graphs + spec requires the MSA path: its verify batches # are decode-shaped and capture-safe (the reference path @@ -7706,52 +7713,55 @@ def test_nvfp4_eagle3(self, tp_size, ep_size, attention_dp, ) if cuda_graph else None, disable_overlap_scheduler=not overlap_scheduler, enable_attention_dp=attention_dp, + enable_iter_perf_stats=True, trust_remote_code=True) as llm: assert llm.args.quant_config.quant_algo == QuantAlgo.MIXED_PRECISION + + def drain_spec_stats(llm): + drafted = accepted = steps = 0 + for s in llm.get_stats(timeout=2): + s = json.loads(s) if isinstance(s, str) else s + sd = s.get("specDecodingStats") or {} + drafted += sd.get("numDraftTokens", 0) + accepted += sd.get("numAcceptedTokens", 0) + steps += sd.get("numRequestsWithDraftTokens", 0) + return drafted, accepted, steps + task = MMLU(model_name) task.evaluate(llm) task = GSM8K(model_name) task.evaluate(llm) - # Acceptance probe (pattern: TestNemotronV3Ultra - # test_nvfp4_4gpu_mtp_ar): stream a few greedy prompts and - # derive per-step acceptance from the token increments. - raw_prompts = [ - "Solve step by step: what is 12 times 17?", - "Write a Python function that reverses a linked list.", - "The capital of France is", - ] - prompts = [ - llm.tokenizer.apply_chat_template( - [{ - "role": "user", - "content": p - }], - tokenize=False, - add_generation_prompt=True, - ) for p in raw_prompts + # 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. + questions = [ + r["question"] + for r in load_dataset("gsm8k", "main", split="test") + ][:200] + chat_prompts = [ + llm.tokenizer.apply_chat_template([{ + "role": "user", + "content": q + }], + tokenize=False, + add_generation_prompt=True) + for q in questions ] - tok_ids = [llm.tokenizer.encode(p) for p in prompts] - sampling_params = SamplingParams(max_tokens=128, temperature=0) - total_drafted = 0 - total_accepted = 0 - total_steps = 0 - for i in range(len(tok_ids)): - num_tokens = 0 - for output in llm.generate_async(tok_ids[i], - sampling_params, - streaming=True): - new_tokens = output.outputs[0].token_ids - total_drafted += max_draft_len - total_accepted += len(new_tokens) - num_tokens - 1 - total_steps += 1 - num_tokens = len(new_tokens) - accept_rate = total_accepted / total_drafted - accept_length = 1 + total_accepted / total_steps - print(f"MiniMax-M3 Eagle3 acceptance: rate={accept_rate:.3f}, " - f"mean acceptance length={accept_length:.3f}") - assert accept_rate > 0.25, \ - f"Eagle3 acceptance rate too low: {accept_rate:.3f}" + drain_spec_stats(llm) + llm.generate(chat_prompts, + SamplingParams(max_tokens=512, temperature=0)) + drafted, accepted, steps = drain_spec_stats(llm) + assert steps > 0, "no speculative iterations recorded" + chat_rate = accepted / drafted + chat_length = 1 + accepted / steps + 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, \ + f"Eagle3 chat-GSM8K acceptance rate too low: {chat_rate:.3f}" + assert chat_length > 3.3, \ + f"Eagle3 chat-GSM8K acceptance length too low: {chat_length:.3f}" @skip_pre_blackwell 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 624ba4b19ff9..1c777d593ae9 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 @@ -224,6 +224,7 @@ def test_msa_proxy_max_score_strided_index_k_matches_packed(): assert index_k_strided.stride(0) == coalescing_scale * page_size * head_dim assert torch.equal(strided_scores, packed_scores) + def test_msa_scratch_sizing_covers_spec_verify_tokens(): """Under one-model Eagle3 spec verify a decode step carries 1 + draft_len query tokens per request, so the proxy scratch must be @@ -252,7 +253,8 @@ def test_per_token_valid_blocks_multi_token_decode(): """Spec-verify decode rows expose one entry per query TOKEN, walking the causal ladder within the verify window.""" from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import ( - per_token_valid_blocks, ) + per_token_valid_blocks, + ) # One request verifying 4 tokens against kv_len 10 (offset 6): token t # attends 7 + t positions; with 2-token blocks that is ceil((7+t)/2).