diff --git a/cpp/tensorrt_llm/common/attentionOp.cpp b/cpp/tensorrt_llm/common/attentionOp.cpp index b91c0ef98df6..3bec69f23ba2 100644 --- a/cpp/tensorrt_llm/common/attentionOp.cpp +++ b/cpp/tensorrt_llm/common/attentionOp.cpp @@ -216,6 +216,7 @@ bool AttentionOp::convertMMHAParamsToXQAParams(tensorrt_llm::kernels::XQAParams& // Medusa mode will have multiple query tokens. xqaParams.multi_query_tokens = mIsSpecDecodingEnabled && mUseSpecDecoding; xqaParams.is_spec_dec_tree = mIsSpecDecTree; + xqaParams.force_prepare_spec_dec_tree_mask = mForcePrepareSpecDecTreeMask; xqaParams.layer_idx = generationsParams.layer_idx; if (mKVCacheQuantMode.hasInt8KvCache()) diff --git a/cpp/tensorrt_llm/common/attentionOp.h b/cpp/tensorrt_llm/common/attentionOp.h index f7337c9c9cb2..c76907803c69 100644 --- a/cpp/tensorrt_llm/common/attentionOp.h +++ b/cpp/tensorrt_llm/common/attentionOp.h @@ -501,6 +501,7 @@ class AttentionOp int32_t mSpecDecodingMaxGenerationLength = 1; // Static spec-dec tree length used by FMHA autotuning. int32_t mSpecDecodingTargetMaxGenLen = 0; + bool mForcePrepareSpecDecTreeMask = false; bool mIsMLAEnabled = false; bool mIsGenerationMLA = false; bool mUseGenFlashMLA = false; @@ -571,11 +572,11 @@ class AttentionOp mCrossAttention, mMaxDistance, mPosShiftEnabled, mPagedContextFMHA, mFP8ContextFMHA, mFP8AttenOutput, mFP8ContextMLA, mFP8GenerationMLA, mChunkPrefillBufferBatchSize, mDenseContextFMHA, mHasFullAttentionMask, mIsSpecDecodingEnabled, mUseSpecDecoding, mIsSpecDecTree, mSpecDecodingIsGenerationLengthVariable, - mSpecDecodingMaxGenerationLength, mSpecDecodingTargetMaxGenLen, mIsMLAEnabled, mIsGenerationMLA, - mUseGenFlashMLA, mUseSparseAttention, mUseTllmGenSparseAttentionPaged, mUseTllmGenSparseAttention, - mMLAParams.data(), mCpSize, mCpRank, mCpGroup, mNumAttnHeads, mNumAttnKVHeads, mNumKVHeadsOrigin, - mAttnTpSize, mAttnTpRank, mAttnCpSize, mAttnCpRank, mUlyssesMQABroadcast, mEnableContextFMHA, - mFMHAForceFP32Acc, mMultiBlockMode, mEnableXQA, mUseKVCache, mSkipAttn, mFuseFp4Quant, + mSpecDecodingMaxGenerationLength, mSpecDecodingTargetMaxGenLen, mForcePrepareSpecDecTreeMask, mIsMLAEnabled, + mIsGenerationMLA, mUseGenFlashMLA, mUseSparseAttention, mUseTllmGenSparseAttentionPaged, + mUseTllmGenSparseAttention, mMLAParams.data(), mCpSize, mCpRank, mCpGroup, mNumAttnHeads, mNumAttnKVHeads, + mNumKVHeadsOrigin, mAttnTpSize, mAttnTpRank, mAttnCpSize, mAttnCpRank, mUlyssesMQABroadcast, + mEnableContextFMHA, mFMHAForceFP32Acc, mMultiBlockMode, mEnableXQA, mUseKVCache, mSkipAttn, mFuseFp4Quant, mFusesDsv4InvRopeFp8Quant, mNbMultiBlockSemaphores, mAttentionChunkSize.value_or(-1), mSkipSoftmaxThresholdScaleFactorPrefill, mSkipSoftmaxThresholdScaleFactorDecode, mSageAttnNumEltsPerBlkQ, mSageAttnNumEltsPerBlkK, mSageAttnNumEltsPerBlkV, mSageAttnQkInt8); diff --git a/cpp/tensorrt_llm/kernels/decoderMaskedMultiheadAttention/xqaParams.h b/cpp/tensorrt_llm/kernels/decoderMaskedMultiheadAttention/xqaParams.h index e421be0a6bd7..dc7794752e47 100644 --- a/cpp/tensorrt_llm/kernels/decoderMaskedMultiheadAttention/xqaParams.h +++ b/cpp/tensorrt_llm/kernels/decoderMaskedMultiheadAttention/xqaParams.h @@ -63,6 +63,7 @@ struct XQAParams int64_t* spec_decoding_bl_tree_mask_offset; // for blackwell spec-dec tree mask offset uint32_t* spec_decoding_bl_tree_mask; // for blackwell spec-dec tree mask int32_t* spec_bl_tree_first_sparse_mask_offset_kv; // for blackwell spec-dec tree first sparse mask offset kv + bool force_prepare_spec_dec_tree_mask = false; int32_t const* mrope_position_deltas = nullptr; // Helix parallelism params. int32_t const* helix_position_offsets = nullptr; diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h index d1681c39a56e..8bcc29e0dbdb 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaKernels.h @@ -496,7 +496,11 @@ class TllmGenFmhaKernel tg::CudaRunner::Grid grid{numCtasX, numCtasY, numCtasZ}; // Prepare custom mask for spec-decoding generation kernels if needed. - if (params.mLayerIdx == 0 && params.mIsSpecDecTree) + bool const prepareSpecDecTreeMask = params.mIsSpecDecTree + && (params.mForcePrepareSpecDecTreeMask || params.mLayerIdx == 0 + || (params.mSpecDecodingTargetMaxGenLen > 0 + && params.mMaxSeqLenQ != params.mSpecDecodingTargetMaxGenLen)); + if (prepareSpecDecTreeMask) { int32_t stepQ = options.mTileSizeQ * options.mNumInstsQ; int32_t stepKv = options.mTileSizeKv * options.mNumInstsKv; diff --git a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunnerParams.h b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunnerParams.h index 83b816c346f5..9252650ac67b 100644 --- a/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunnerParams.h +++ b/cpp/tensorrt_llm/kernels/trtllmGenKernels/fmha/fmhaRunnerParams.h @@ -369,6 +369,7 @@ struct TllmGenFmhaRunnerParams // row stride ceilDiv(mPackedMaskMaxSeqLenQ, 32) rather than ceilDiv(seqLenQ, 32). int32_t mPackedMaskMaxSeqLenQ = 0; int32_t mSpecDecodingTargetMaxGenLen = 0; + bool mForcePrepareSpecDecTreeMask = false; // set the attention mask type TllmGenFmhaRunnerParams& setAttentionMaskType(std::int8_t maskType) diff --git a/cpp/tensorrt_llm/kernels/xqaDispatcher.cpp b/cpp/tensorrt_llm/kernels/xqaDispatcher.cpp index 35fd02e7f127..8a5eeee91c6b 100644 --- a/cpp/tensorrt_llm/kernels/xqaDispatcher.cpp +++ b/cpp/tensorrt_llm/kernels/xqaDispatcher.cpp @@ -534,6 +534,7 @@ void XqaDispatcher::runImpl( tllmRunnerParams.generalPackedCustoMaskPtr = params.spec_decoding_packed_mask; tllmRunnerParams.mPackedMaskMaxSeqLenQ = params.spec_decoding_max_generation_length; tllmRunnerParams.mSpecDecodingTargetMaxGenLen = mFixedParams.specDecodingTargetMaxGenLen; + tllmRunnerParams.mForcePrepareSpecDecTreeMask = params.force_prepare_spec_dec_tree_mask; tllmRunnerParams.customMaskPtr = params.spec_decoding_bl_tree_mask; tllmRunnerParams.customMaskOffsetsPtr = params.spec_decoding_bl_tree_mask_offset; tllmRunnerParams.firstSparseMaskOffsetsKvPtr = params.spec_bl_tree_first_sparse_mask_offset_kv; diff --git a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp index 23bd36d67b82..dd1ff0db5410 100644 --- a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp @@ -179,7 +179,8 @@ void initBindings(nb::module_& m) nb::arg("relative_attention_bias") = std::nullopt, nb::arg("relative_attention_max_distance") = 0, nb::arg("spec_decoding_target_max_draft_tokens") = std::nullopt, nb::arg("quant_scale_qkv") = std::nullopt, nb::arg("dsv4_inv_rope_cos_sin_cache") = std::nullopt, nb::arg("enable_dsv4_epilogue_fusion") = false, - "Multi-head attention operation", nb::call_guard()); + nb::arg("force_prepare_spec_dec_tree_mask") = false, "Multi-head attention operation", + nb::call_guard()); m.def( "get_helix_workspace_size_per_rank", diff --git a/cpp/tensorrt_llm/thop/attentionOp.cpp b/cpp/tensorrt_llm/thop/attentionOp.cpp index 848554c512d6..0bf5bc001206 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.cpp +++ b/cpp/tensorrt_llm/thop/attentionOp.cpp @@ -1084,7 +1084,8 @@ void attention(torch::Tensor q, std::optional k, std::optional compressed_kv_cache_pool_ptr, bool const is_cross, std::optional cross_kv, std::optional relative_attention_bias, int64_t relative_attention_max_distance, std::optional spec_decoding_target_max_draft_tokens, std::optional quant_scale_qkv, - std::optional dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion) + std::optional dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion, + bool const force_prepare_spec_dec_tree_mask) { TLLM_LOG_TRACE("Attention op starts at layer %d", local_layer_idx); // Use these tensors to infer if the attention is using KV cache @@ -1237,6 +1238,7 @@ void attention(torch::Tensor q, std::optional k, std::optionalmSpecDecodingTargetMaxGenLen = static_cast(spec_decoding_target_max_draft_tokens.value()) + 1; } + op->mForcePrepareSpecDecTreeMask = force_prepare_spec_dec_tree_mask; op->mUseSparseAttention = false; op->mUseTllmGenSparseAttentionPaged = false; diff --git a/cpp/tensorrt_llm/thop/attentionOp.h b/cpp/tensorrt_llm/thop/attentionOp.h index b4f0e90bc9a7..f28209166e0d 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.h +++ b/cpp/tensorrt_llm/thop/attentionOp.h @@ -95,7 +95,8 @@ void attention(torch::Tensor q, std::optional k, std::optional relative_attention_bias = std::nullopt, int64_t relative_attention_max_distance = 0, std::optional spec_decoding_target_max_draft_tokens = std::nullopt, std::optional quant_scale_qkv = std::nullopt, - std::optional dsv4_inv_rope_cos_sin_cache = std::nullopt, bool enable_dsv4_epilogue_fusion = false); + std::optional dsv4_inv_rope_cos_sin_cache = std::nullopt, bool enable_dsv4_epilogue_fusion = false, + bool const force_prepare_spec_dec_tree_mask = false); struct KvCachePoolPointers { diff --git a/docs/source/features/feature-combination-matrix.md b/docs/source/features/feature-combination-matrix.md index b56c1b219af9..d91322fc8c11 100644 --- a/docs/source/features/feature-combination-matrix.md +++ b/docs/source/features/feature-combination-matrix.md @@ -14,7 +14,7 @@ | Speculative Decoding — Linear | Yes | Yes | Yes | No | Yes | No | Yes | Yes | Yes | --- | | | | | | | | | | | Speculative Decoding — Dynamic Trees | Yes | Yes | Yes | No | Yes | No | Yes | Yes | Yes | No | --- | | | | | | | | | | Speculative Decoding — Legacy Path (NGram, user-provided) | Yes | Yes | Yes | No | Yes | No | Yes | Yes | Yes | No | No | --- | | | | | | | | -| Torch Sampler | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | --- | | | | | | | +| Torch Sampler | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes (MTP dynamic tree: greedy only) | Yes | --- | | | | | | | | TLLM C++ Sampler | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | No | No | No | No | --- | | | | | | | KV Cache Reuse | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | --- | | | | | | Sliding Window Attention | Yes | Yes | Yes | Yes | Yes | Untested | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | Yes | --- | | | | diff --git a/examples/llm-api/quickstart_advanced.py b/examples/llm-api/quickstart_advanced.py index 718d4a410e2c..d2acd490d44a 100644 --- a/examples/llm-api/quickstart_advanced.py +++ b/examples/llm-api/quickstart_advanced.py @@ -310,6 +310,9 @@ def setup_llm(args, **kwargs): relaxed_topk=args.relaxed_topk, relaxed_delta=args.relaxed_delta, mtp_eagle_one_model=args.use_one_model, + use_dynamic_tree=args.use_dynamic_tree, + dynamic_tree_max_topK=args.dynamic_tree_max_topK, + max_total_draft_tokens=args.max_total_draft_tokens, speculative_model=args.model_dir) elif spec_decode_algo == "EAGLE3": spec_config = Eagle3DecodingConfig( diff --git a/tensorrt_llm/_torch/attention_backend/fmha/fallback.py b/tensorrt_llm/_torch/attention_backend/fmha/fallback.py index 8f1da5dc36f9..9c081bf19dfc 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/fallback.py @@ -98,6 +98,7 @@ def forward( spec_decoding_bl_tree_mask_offset=metadata.spec_decoding_bl_tree_mask_offset, spec_decoding_bl_tree_mask=metadata.spec_decoding_bl_tree_mask, spec_decoding_target_max_draft_tokens=metadata.max_total_draft_tokens, + force_prepare_spec_dec_tree_mask=metadata.force_prepare_spec_dec_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv=metadata.spec_bl_tree_first_sparse_mask_offset_kv, num_sparse_topk=metadata.num_sparse_topk, flash_mla_tile_scheduler_metadata=metadata.flash_mla_tile_scheduler_metadata, diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 40d15970398b..a3f4e89d8d85 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -117,6 +117,7 @@ def effective_beam_width(self) -> int: is_spec_dec_tree: bool = False # if spec-dec tree wouldn't be changed at all, the mask won't be computed every step. is_spec_dec_dynamic_tree: bool = False + force_prepare_spec_dec_tree_mask: bool = False # parameters required for spec-dec mode max_total_draft_tokens: Optional[int] = None diff --git a/tensorrt_llm/_torch/models/modeling_nemotron_h.py b/tensorrt_llm/_torch/models/modeling_nemotron_h.py index 5c2e60c4706a..3bb491fd1b91 100644 --- a/tensorrt_llm/_torch/models/modeling_nemotron_h.py +++ b/tensorrt_llm/_torch/models/modeling_nemotron_h.py @@ -1275,6 +1275,8 @@ def forward( residual=residual, attn_metadata=attn_metadata, all_rank_num_tokens=all_rank_num_tokens, + spec_metadata=spec_metadata, + mamba_metadata=attn_metadata.mamba_metadata, lora_params=lora_params, ) return hidden_states diff --git a/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py b/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py index e807a904978f..688c48b147de 100644 --- a/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py +++ b/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py @@ -430,16 +430,47 @@ def forward( f"{draft_token_num} must match fixed replay step " f"width {replay_step_width}.") - intermediate_state_indices = _cached_arange( - attn_metadata.kv_cache_manager.get_max_resource_count(), - state_indices_d.device)[:num_decodes] + # Dynamic-tree verify uses per-request links; linear MTP skips it. + is_dyn_tree = getattr(spec_metadata, 'is_spec_dec_dynamic_tree', + False) + retrieve_next_token = retrieve_next_sibling = None + retrieve_parent_token = None + if is_dyn_tree: + if use_replay: + raise NotImplementedError( + "Dynamic-tree Mamba verify is not supported with " + "the replay SSM-cache path (TRTLLM_USE_MAMBA_REPLAY)." + ) + retrieve_next_token = spec_metadata.retrieve_next_token + retrieve_next_sibling = spec_metadata.retrieve_next_sibling + assert (retrieve_next_token is not None + and retrieve_next_sibling is not None), ( + "Dynamic-tree verify requires retrieve link " + "tensors on spec_metadata.") + retrieve_next_token = retrieve_next_token[:num_decodes] + retrieve_next_sibling = retrieve_next_sibling[:num_decodes] + # conv1d fills parent links used by tree-aware SSM restore. + retrieve_parent_token = torch.empty( + (num_decodes, draft_token_num), + dtype=torch.int32, + device=state_indices_d.device) + + # Prefer the cache_manager-owned arange; cached fallback storage + # can be recycled by CUDA graph warmup. + _km_isi = getattr(attn_metadata.kv_cache_manager, + 'intermediate_state_indices', None) + if _km_isi is not None: + intermediate_state_indices = _km_isi[:num_decodes] + else: + intermediate_state_indices = _cached_arange( + attn_metadata.kv_cache_manager.get_max_resource_count(), + state_indices_d.device)[:num_decodes] - # Reshape for batch processing - xbc_d_reshaped = xbc_d.view(num_decodes, draft_token_num, - -1).transpose(1, 2) + # Use reshape because dynamic-tree tokens may be non-contiguous. + xbc_d_reshaped = xbc_d.reshape(num_decodes, draft_token_num, + -1).transpose(1, 2) def conv1d(): - # TODO:support tree structure [TRTLLM-10320] xbc_d_processed = causal_conv1d_update_triton( xbc_d_reshaped, conv_states, @@ -449,11 +480,15 @@ def conv1d(): conv_state_indices=state_indices_d[:num_decodes], intermediate_conv_window=intermediate_conv_states, intermediate_state_indices=intermediate_state_indices, + # None on linear MTP. + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_parent_token=retrieve_parent_token, # PDL chain: conv1d → precompute → main (replay only) launch_dependent_kernels=use_replay, ) - return xbc_d_processed.transpose(1, 2).view( + return xbc_d_processed.transpose(1, 2).reshape( num_decode_tokens, -1) else: @@ -588,13 +623,23 @@ def convert_dt(): state_batch_indices=state_batch_indices, disable_state_update=True, intermediate_state_indices=intermediate_state_indices, + # None for linear MTP; tree parent map for dynamic tree. + retrieve_parent_token=retrieve_parent_token, ) else: # Triton kernel + flashinfer need contiguous for alignment. x_d_4d = x_d_4d.contiguous() B_d_4d = B_d_4d.contiguous() C_d_4d = C_d_4d.contiguous() - self.selective_state_update_func( + if is_dyn_tree: + # flashinfer SSU cannot restore tree-parent states. + ssu_func = selective_state_update_native + ssu_extra = dict( + retrieve_parent_token=retrieve_parent_token) + else: + ssu_func = self.selective_state_update_func + ssu_extra = {} + ssu_func( ssm_states, x_d_4d, dt_d_4d, @@ -611,6 +656,7 @@ def convert_dt(): intermediate_states_buffer=intermediate_ssm_states, cache_steps=draft_token_num, intermediate_state_indices=intermediate_state_indices, + **ssu_extra, **philox_kwargs, ) else: diff --git a/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py b/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py index d02e19d5b95c..0274ac32ceae 100644 --- a/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py @@ -784,14 +784,20 @@ def _drop(tensor): torch.cuda.empty_cache() @torch.compile(options={"max-autotune": True}) - def update_mamba_states(self, attn_metadata: "AttentionMetadata", - num_accepted_tokens: torch.Tensor, - state_indices: torch.Tensor): + def update_mamba_states( + self, + attn_metadata: "AttentionMetadata", + num_accepted_tokens: torch.Tensor, + state_indices: torch.Tensor, + accepted_leaf_positions: Optional[torch.Tensor] = None): batch_size = attn_metadata.num_seqs num_contexts = attn_metadata.num_contexts num_gens = batch_size - num_contexts num_accepted_draft_tokens = num_accepted_tokens[ num_contexts:num_contexts + num_gens] - 1 + # Dynamic tree passes tree-node leaf positions; linear MTP uses depth. + accepted_positions = (accepted_leaf_positions if accepted_leaf_positions + is not None else num_accepted_draft_tokens) state_indices_d = state_indices[num_contexts:num_contexts + num_gens] src_state_indices = self.intermediate_state_indices[:num_gens] @@ -828,7 +834,7 @@ def update_mamba_states(self, attn_metadata: "AttentionMetadata", ssm_states = self.mamba_cache.temporal intermediate_ssm_cache = self.mamba_cache.intermediate_ssm accepted_ssm_state = intermediate_ssm_cache[:, src_state_indices, - num_accepted_draft_tokens] + accepted_positions] ssm_states[:, state_indices_d, :] = accepted_ssm_state # Conv: both paths save all intermediate conv windows, carry over the accepted one. @@ -836,7 +842,7 @@ def update_mamba_states(self, attn_metadata: "AttentionMetadata", intermediate_conv_window_cache = self.mamba_cache.intermediate_conv_window accepted_conv_state = intermediate_conv_window_cache[:, src_state_indices, - num_accepted_draft_tokens] + accepted_positions] conv_states[:, state_indices_d, :] = accepted_conv_state @@ -980,15 +986,18 @@ def mamba_layer_cache( def shutdown(self): self._impl.shutdown() - def update_mamba_states(self, attn_metadata: "AttentionMetadata", - num_accepted_tokens: torch.Tensor, - state_indices: torch.Tensor): + def update_mamba_states( + self, + attn_metadata: "AttentionMetadata", + num_accepted_tokens: torch.Tensor, + state_indices: torch.Tensor, + accepted_leaf_positions: Optional[torch.Tensor] = None): # Non-speculative configs don't allocate intermediate state; the # promotion is a clean no-op. if not self._impl.is_speculative(): return self._impl.update_mamba_states(attn_metadata, num_accepted_tokens, - state_indices) + state_indices, accepted_leaf_positions) class MambaHybridCacheManager(BaseResourceManager, BaseMambaCacheManager): @@ -1755,10 +1764,12 @@ def is_speculative(self) -> bool: return self.spec_config is not None @nvtx_range("hybrid_update_mamba_states") - def update_mamba_states(self, - attn_metadata: "AttentionMetadata", - num_accepted_tokens: torch.Tensor, - state_indices: Optional[torch.Tensor] = None): + def update_mamba_states( + self, + attn_metadata: "AttentionMetadata", + num_accepted_tokens: torch.Tensor, + state_indices: Optional[torch.Tensor] = None, + accepted_leaf_positions: Optional[torch.Tensor] = None): if self.local_num_mamba_layers == 0: return batch_size = attn_metadata.num_seqs @@ -1767,6 +1778,10 @@ def update_mamba_states(self, num_accepted_draft_tokens = ( num_accepted_tokens[num_contexts:num_contexts + num_gens] - 1).to( torch.int32) + # Dynamic tree passes tree-node leaf positions; linear MTP uses depth. + accepted_positions = (accepted_leaf_positions.to(torch.int32) + if accepted_leaf_positions is not None else + num_accepted_draft_tokens) # Match the API of MambaCacheManager.update_mamba_states: callers # may pass per-request state slot indices explicitly (e.g. MTP via # attn_metadata.mamba_metadata.state_indices). Fall back to this @@ -1814,16 +1829,15 @@ def update_mamba_states(self, # Legacy: copy the accepted SSM state from the intermediate buffer. _promote_mamba_state_triton(self.all_ssm_states, self.intermediate_ssm_states, - src_state_indices, - num_accepted_draft_tokens, + src_state_indices, accepted_positions, state_indices_d) # Conv: both paths save all intermediate conv windows, carry over the # accepted one. _promote_mamba_state_triton(self.all_conv_states, self.intermediate_conv_states, - src_state_indices, - num_accepted_draft_tokens, state_indices_d) + src_state_indices, accepted_positions, + state_indices_d) @torch.inference_mode() def _refresh_dummy_request_mask(self, is_dummy: List[bool]) -> None: diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index d83f51a1bbe3..6adfc7ec14cc 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -620,7 +620,13 @@ def __init__( ) or self.model_is_wrapped self.max_total_draft_tokens = spec_config.tokens_per_gen_step - 1 self.max_draft_len = spec_config.max_draft_len - self.runtime_draft_len = spec_config.max_draft_len + # Mutable per-iteration draft length (updated each iteration when + # dynamic draft length is enabled; otherwise stays fixed). Tree + # modes verify all tree nodes per step, which can be wider than the + # tree depth used by the drafter loop. + self.runtime_draft_len = (self.max_total_draft_tokens + if not spec_config.is_linear_tree else + self.max_draft_len) else: self.without_logits = False diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index d9b5be1ecc1e..062b61c03be9 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -398,6 +398,10 @@ class Eagle3OneModelSpecMetadata(SpecMetadata): # prepare() before self.num_tokens is decremented to the attention-DP subseq # shape; maybe_capture_hidden_states must bound by this, not self.num_tokens. num_capture_tokens: int = 0 + # Per-generation tree links for Mamba verify in dynamic-tree one-model paths. + retrieve_next_token: Optional[torch.Tensor] = None + retrieve_next_sibling: Optional[torch.Tensor] = None + retrieve_parent_token: Optional[torch.Tensor] = None def __post_init__(self): if self.layers_to_capture is None: @@ -443,9 +447,10 @@ def __post_init__(self): self.hidden_size * len(self.layers_to_capture)), dtype=self.dtype, device='cuda') - if (self.spec_resource_manager is not None - and self.spec_resource_manager.batch_indices_cuda is not None): - self.batch_indices_cuda = self.spec_resource_manager.batch_indices_cuda + batch_indices_cuda = getattr(self.spec_resource_manager, + "batch_indices_cuda", None) + if batch_indices_cuda is not None: + self.batch_indices_cuda = batch_indices_cuda assert self.batch_indices_cuda.shape[0] >= self.max_num_requests, ( f"batch_indices_cuda shape mismatch: " f"{type(self.spec_resource_manager).__name__} has " @@ -530,6 +535,23 @@ def prepare(self): if gen_request_ids: sa_manager.prepare(gen_request_ids, self.runtime_draft_len) + self.retrieve_next_token = None + self.retrieve_next_sibling = None + self.retrieve_parent_token = None + spec_tree_manager = getattr(self.spec_resource_manager, + 'spec_tree_manager', None) + if self.use_dynamic_tree and spec_tree_manager is not None: + num_gens = self.num_generations + if num_gens > 0: + num_contexts = num_seqs - num_gens + slot_storage = spec_tree_manager.slot_storage + gen_slot_ids = slot_storage.all_ids_buf[ + num_contexts:num_contexts + num_gens] + next_token, next_sibling = slot_storage.next_links_from_slots( + gen_slot_ids, num_gens) + self.retrieve_next_token = next_token + self.retrieve_next_sibling = next_sibling + def maybe_capture_hidden_states( self, layer_id: int, diff --git a/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py new file mode 100644 index 000000000000..5dca84abae8e --- /dev/null +++ b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py @@ -0,0 +1,1122 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""MTP-Eagle one-model dynamic tree speculative decoding (greedy only).""" + +import math +from typing import TYPE_CHECKING, List, Optional + +import torch +import triton + +from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import MambaHybridCacheManager +from tensorrt_llm._utils import get_sm_version, nvtx_range +from tensorrt_llm.mapping import Mapping + +from ..distributed.ops import allgather +from ..model_config import ModelConfig +from ..pyexecutor.llm_request import LlmRequest +from ..pyexecutor.resource_manager import BaseResourceManager +from ..pyexecutor.scheduler import ScheduledRequests +from .eagle3 import MTPEagleWorker + +# Reuse drafter-agnostic dynamic-tree helpers. +from .eagle3_dynamic_tree import ( + _build_mask_and_position, + _gather_repack_step0_kernel, + _resample_final_tokens, + _select_topk_draft_tokens, +) +from .mtp import MTPHiddenStatesManager + +if TYPE_CHECKING: + from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig + + +class MTPEagleDynamicTreeWorker(MTPEagleWorker): + """MTP-Eagle worker with dynamic-tree draft and greedy verify.""" + + def __init__( + self, + spec_config: "MTPDecodingConfig", + model_config: Optional[ModelConfig] = None, + use_separate_draft_kv_cache: bool = False, + *, + mapping: Optional[Mapping] = None, + ): + super().__init__(spec_config, model_config, use_separate_draft_kv_cache, mapping=mapping) + assert getattr(spec_config, "use_dynamic_tree", False), ( + "MTPEagleDynamicTreeWorker requires use_dynamic_tree=True" + ) + + from .dynamic_tree_ops import DynamicTreeOpsConverter + + self.K = spec_config.dynamic_tree_max_topK + self.max_total_draft_tokens = spec_config.tokens_per_gen_step - 1 + self.tokens_per_gen_step = spec_config.tokens_per_gen_step + # Set by py_executor_creator from the global max_batch_size. + assert spec_config._max_batch_size is not None, ( + "MTPDecodingConfig._max_batch_size was not populated; " + "py_executor_creator should have set it from the global max_batch_size." + ) + self._max_batch_size = spec_config._max_batch_size + + K = self.K + max_draft_len = spec_config.max_draft_len + max_batch_size = self._max_batch_size + loop_max_tokens = K * max_draft_len # draft loop working size + + # spec_tree_manager is lazily bound from the resource manager. + self.spec_tree_manager = None + self._d2t = None + + # Pre-allocated draft-loop buffers (CUDA-graph safe). + self.draft_tokens_buffer = torch.zeros( + max_batch_size, loop_max_tokens, dtype=torch.int32, device="cuda" + ) + self.position_ids_buffer = torch.zeros( + max_batch_size, loop_max_tokens, dtype=torch.int32, device="cuda" + ) + self.history_draft_tokens_buffer = torch.zeros( + (max_batch_size, (K + K * K * (max_draft_len - 1))), dtype=torch.int32, device="cuda" + ) + self.history_score_buffer = torch.zeros( + (max_batch_size, K + K * K * (max_draft_len - 1)), dtype=torch.float32, device="cuda" + ) + self.history_draft_tokens_parent_buffer = torch.zeros( + (max_batch_size, max(K * (max_draft_len - 1) + 1, K + 1)), + dtype=torch.int64, + device="cuda", + ) + self.tree_mask_buffer = torch.zeros( + (max_batch_size * loop_max_tokens * loop_max_tokens), dtype=torch.int32, device="cuda" + ) + self.tree_mask_init_buffer = ( + torch.eye(K, dtype=torch.int32, device="cuda").unsqueeze(0).repeat(max_batch_size, 1, 1) + ) + self.tree_ops_converter = DynamicTreeOpsConverter( + dynamic_tree_max_topK=K, + max_draft_len=max_draft_len, + max_total_draft_tokens=self.max_total_draft_tokens, + max_batch_size=max_batch_size, + device=torch.device("cuda"), + ) + + self._max_path_len = max_draft_len + 1 + # Step-0 draft resets verify-time tree metadata to accepted-path width. + self._kv_correction = self.tokens_per_gen_step - self._max_path_len + self._step0_causal_mask = torch.tensor( + [(1 << (t + 1)) - 1 for t in range(self._max_path_len)], + dtype=torch.int32, + device="cuda", + ) + self._causal_offs = torch.arange(self._max_path_len, device="cuda", dtype=torch.int32) + self._last_selected_parents = None + self._parent_init_arange = torch.arange(-1, K, device="cuda", dtype=torch.int32) + + # Accepted-path bookkeeping for KV relocation and output. + self._accepted_draft_indices_tensor = torch.full( + (max_batch_size, max_draft_len), -1, dtype=torch.int32, device="cuda" + ) + self._kv_head_dim_bytes = None + + # === Verification buffers (greedy only) === + N = self.tokens_per_gen_step + self._accepted_tokens_buf = torch.zeros( + max_batch_size, self._max_path_len, dtype=torch.int32, device="cuda" + ) + self._num_accepted_tokens_buf = torch.ones(max_batch_size, dtype=torch.int32, device="cuda") + self._target_tokens_buf = torch.zeros(max_batch_size * N, dtype=torch.int64, device="cuda") + self._candidates_buf = torch.zeros(max_batch_size, N, dtype=torch.int32, device="cuda") + self._target_predict_buf = torch.zeros(max_batch_size, N, dtype=torch.int32, device="cuda") + + # Hidden states for the growing-context draft loop. + self._hs_write_buffer = None + self._accumulated_hs = None + self._hs_read_map = torch.zeros( + max_batch_size, loop_max_tokens, dtype=torch.long, device="cuda" + ) + self._step0_hs = None + self._hs_dim = None + + # Step-0 repack scratch for accepted-path inputs. + max_total_tokens = max_batch_size * self.tokens_per_gen_step + self._step0_input_ids_buf = torch.zeros(max_total_tokens, dtype=torch.int32, device="cuda") + self._step0_position_ids_buf = torch.zeros( + max_total_tokens, dtype=torch.int32, device="cuda" + ) + self._step0_hidden_states_buf = None + self._gather_ids_buf = torch.zeros(max_total_tokens, dtype=torch.long, device="cuda") + + # Mask repack scratch (graph-safe; avoids .contiguous() in the loop). + buf_dim = max(self.max_total_draft_tokens + 1, loop_max_tokens) + mask_width = (buf_dim + 31) // 32 + self._mask_repack_buf = torch.zeros( + max_batch_size * buf_dim * mask_width, dtype=torch.int32, device="cuda" + ) + # sm>=100 (except 120/121): prepareCustomMask keeps padded 3D; no repack. + sm = get_sm_version() + self._needs_mask_repack = sm < 100 or sm in (120, 121) + + def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): + super()._prepare_attn_metadata_for_spec_dec(attn_metadata) + + batch_size = attn_metadata.num_seqs + if hasattr(attn_metadata, "kv_lens_cuda"): + # Keep kv_lens_cuda itself alive because TRTLLM attention holds a + # runtime view into it. + self._saved_kv_lens_cuda = attn_metadata.kv_lens_cuda[:batch_size].clone() + else: + self._saved_kv_lens_cuda = None + + # Restore verify metadata after the draft loop mutates it. + if attn_metadata.spec_decoding_packed_mask is not None: + self._saved_packed_mask = attn_metadata.spec_decoding_packed_mask[:batch_size].clone() + else: + self._saved_packed_mask = None + if attn_metadata.spec_decoding_position_offsets is not None: + self._saved_position_offsets = attn_metadata.spec_decoding_position_offsets.clone() + self._saved_position_offsets_cpp = attn_metadata.spec_decoding_position_offsets_cpp + else: + self._saved_position_offsets = None + self._saved_position_offsets_cpp = None + if attn_metadata.spec_decoding_generation_lengths is not None: + self._saved_generation_lengths = attn_metadata.spec_decoding_generation_lengths[ + :batch_size + ].clone() + else: + self._saved_generation_lengths = None + + def prepare_position_ids_and_last_tokens(self, position_ids, seq_lens_cuda): + position_ids = position_ids.squeeze(0) + last_tokens_idx = torch.cumsum(seq_lens_cuda, dim=0, dtype=torch.long) - 1 + return position_ids, last_tokens_idx + + def _restore_attn_metadata_from_spec_dec(self, attn_metadata): + super()._restore_attn_metadata_from_spec_dec(attn_metadata) + + if self._saved_kv_lens_cuda is not None: + batch_size = self._saved_kv_lens_cuda.shape[0] + attn_metadata.kv_lens_cuda[:batch_size].copy_(self._saved_kv_lens_cuda) + self._saved_kv_lens_cuda = None + + if self._saved_packed_mask is not None: + batch_size = self._saved_packed_mask.shape[0] + attn_metadata.spec_decoding_packed_mask[:batch_size].copy_(self._saved_packed_mask) + self._saved_packed_mask = None + if self._saved_position_offsets is not None: + attn_metadata.spec_decoding_position_offsets.copy_(self._saved_position_offsets) + attn_metadata.spec_decoding_position_offsets_cpp = self._saved_position_offsets_cpp + self._saved_position_offsets = None + self._saved_position_offsets_cpp = None + if self._saved_generation_lengths is not None: + batch_size = self._saved_generation_lengths.shape[0] + attn_metadata.spec_decoding_generation_lengths[:batch_size].copy_( + self._saved_generation_lengths + ) + self._saved_generation_lengths = None + + # ------------------------------------------------------------------ # + # Helpers # + # ------------------------------------------------------------------ # + def _apply_spec_metadata(self, attn_metadata, batch_size, query_len): + """Set spec-dec gen lengths and refresh the C++ position-offset view.""" + attn_metadata.spec_decoding_generation_lengths[:batch_size] = query_len + attn_metadata.update_position_offsets_for_cpp(query_len) + + def _refresh_blackwell_tree_mask_metadata(self, attn_metadata): + if not getattr(attn_metadata, "use_spec_decoding", False): + return + if not getattr(attn_metadata, "is_spec_dec_dynamic_tree", False): + return + + first_sparse = getattr(attn_metadata, "spec_bl_tree_first_sparse_mask_offset_kv", None) + bl_tree_mask = getattr(attn_metadata, "spec_decoding_bl_tree_mask", None) + if first_sparse is None and bl_tree_mask is None: + return + + if bl_tree_mask is not None: + bl_tree_mask.zero_() + if first_sparse is not None: + attn_metadata.update_blackwell_first_sparse_mask_offset() + + def _repack_mask_padded_to_packed(self, mask_buf, n_req, n_tok): + """Compact padded masks into the flat prefix XQA expects.""" + buf_dim = mask_buf.shape[1] + if n_tok >= buf_dim or n_req <= 1: + return + mask_width = math.ceil(n_tok / 32) + total_elems = n_req * n_tok * mask_width + scratch = self._mask_repack_buf[:total_elems].view(n_req, n_tok, mask_width) + scratch.copy_(mask_buf[:n_req, :n_tok, :mask_width]) + flat = mask_buf.view(-1) + flat[:total_elems] = scratch.view(-1) + + @nvtx_range("mtp_dyn._ensure_spec_tree_manager") + def _ensure_spec_tree_manager(self, resource_manager): + """Lazily bind spec_tree_manager and KV head metadata.""" + if self.spec_tree_manager is not None: + return + from ..pyexecutor.resource_manager import ResourceManagerType + + spec_rm = resource_manager.get_resource_manager(ResourceManagerType.SPEC_RESOURCE_MANAGER) + assert spec_rm is not None and hasattr(spec_rm, "spec_tree_manager"), ( + "Dynamic tree mode requires spec_tree_manager in resource_manager" + ) + self.spec_tree_manager = spec_rm.spec_tree_manager + + if self._kv_head_dim_bytes is None: + cache_mgr = resource_manager.get_resource_manager(ResourceManagerType.KV_CACHE_MANAGER) + if cache_mgr is not None and hasattr(cache_mgr, "head_dim"): + from tensorrt_llm.bindings import DataType + + _dtype_bytes = { + DataType.HALF: 2, + DataType.BF16: 2, + DataType.FLOAT: 4, + DataType.FP8: 1, + DataType.INT8: 1, + DataType.NVFP4: 0.5, + } + self._kv_head_dim_bytes = int( + cache_mgr.head_dim * _dtype_bytes.get(cache_mgr.dtype, 0.5) + ) + + @nvtx_range("mtp_dyn.sample") + def sample( + self, logits: torch.Tensor, max_top_k: int, draft_model=None + ) -> tuple[torch.Tensor, torch.Tensor]: + """TopK sampling for dynamic tree; all-gather sharded TP logits.""" + mapping = ( + getattr(self.model_config, "mapping", None) if self.model_config is not None else None + ) + if mapping is not None and mapping.tp_size > 1 and not mapping.enable_attention_dp: + logits = allgather(logits, mapping, dim=-1) + if draft_model is not None: + vocab_size = draft_model.lm_head.num_embeddings + logits = logits[..., :vocab_size] + probs = torch.softmax(logits, dim=-1) + topk_values, topk_indices = torch.topk(probs, k=max_top_k, dim=-1) + return topk_indices, topk_values + + def update_draft_tokens_and_scores( + self, + cur_draft_idx, + new_draft_tokens, + new_draft_scores, + previous_draft_scores, + batch_size, + attn_metadata=None, + ): + """Grow the tree and update history buffers.""" + if cur_draft_idx == 0: + new_draft_scores = new_draft_scores.reshape(batch_size, self.K) + new_draft_tokens_2d = new_draft_tokens.reshape(batch_size, self.K) + self.draft_tokens_buffer[:batch_size, : self.K] = new_draft_tokens_2d + self.history_draft_tokens_buffer[:batch_size, : self.K] = new_draft_tokens_2d + self.history_score_buffer[:batch_size, : self.K] = new_draft_scores + # Parent buffer: -1 for root, 0..K-1 for first layer. + self.history_draft_tokens_parent_buffer[:batch_size, : self.K + 1] = ( + self._parent_init_arange + ) + self.prepare_tree_mask_and_position_offset(cur_draft_idx, attn_metadata, None) + return new_draft_scores + + ( + real_draft_tokens, + topk_values, + topk_indices, + selected_parents, + new_draft_tokens, + new_draft_scores, + ) = _select_topk_draft_tokens( + new_draft_tokens, new_draft_scores, previous_draft_scores, self.K + ) + + num_tokens_previous_layer = cur_draft_idx * self.K + num_tokens_current_layer = (cur_draft_idx + 1) * self.K + self.draft_tokens_buffer[ + :batch_size, num_tokens_previous_layer:num_tokens_current_layer + ] = real_draft_tokens + + write_start = self.K + (cur_draft_idx - 1) * self.K * self.K + write_end = write_start + self.K * self.K + self.history_draft_tokens_buffer[:batch_size, write_start:write_end] = new_draft_tokens + self.history_score_buffer[:batch_size, write_start:write_end] = new_draft_scores + + self._last_selected_parents = selected_parents + self.prepare_tree_mask_and_position_offset(cur_draft_idx, attn_metadata, selected_parents) + + if cur_draft_idx < self.max_draft_len - 1: + next_layer_start = cur_draft_idx * self.K + 1 + next_layer_end = next_layer_start + self.K + parents_relative_indices = topk_indices + self.K**2 * (cur_draft_idx - 1) + self.K + self.history_draft_tokens_parent_buffer[ + :batch_size, next_layer_start:next_layer_end + ] = parents_relative_indices + return topk_values + + def resampling_final_draft_tokens(self, batch_size: int): + """Reconstruct the final tree from history buffers.""" + return _resample_final_tokens( + self.history_score_buffer[:batch_size, :], + self.history_draft_tokens_buffer[:batch_size, :], + self.max_total_draft_tokens, + ) + + def prepare_tree_mask_and_position_offset( + self, cur_draft_idx, attn_metadata, selected_parents=None + ): + """Prepare mask and position offsets for the next draft layer.""" + if attn_metadata.spec_decoding_packed_mask is None: + return + spec_tree_manager = self.spec_tree_manager + batch_size = attn_metadata.num_seqs + num_tokens_current_layer = self.K * (cur_draft_idx + 1) + num_tokens_previous_layer = self.K * cur_draft_idx + packed_mask = attn_metadata.spec_decoding_packed_mask + if cur_draft_idx == 0: + spec_tree_manager.compute_spec_dec_packed_mask( + self.tree_mask_init_buffer[:batch_size], + packed_mask[:batch_size, :num_tokens_current_layer, :], + ) + self.tree_mask_buffer[ + : batch_size * num_tokens_current_layer * num_tokens_current_layer + ].copy_(self.tree_mask_init_buffer[:batch_size].view(-1)) + attn_metadata.spec_decoding_position_offsets.fill_(0) + self._apply_spec_metadata(attn_metadata, batch_size, num_tokens_current_layer) + else: + num_parent_mask = batch_size * cur_draft_idx * self.K * cur_draft_idx * self.K + parent_mask = self.tree_mask_buffer[:num_parent_mask].reshape( + batch_size, cur_draft_idx * self.K, cur_draft_idx * self.K + ) + + prev_total = batch_size * num_tokens_previous_layer + previous_position_offsets = attn_metadata.spec_decoding_position_offsets[ + :prev_total + ].view(batch_size, num_tokens_previous_layer) + + current_mask, new_positions = _build_mask_and_position( + parent_mask, + selected_parents, + self.tree_mask_init_buffer[:batch_size], + previous_position_offsets, + self.K, + ) + + spec_tree_manager.compute_spec_dec_packed_mask( + current_mask, packed_mask[:batch_size, :num_tokens_current_layer, :] + ) + self.tree_mask_buffer[ + : batch_size * num_tokens_current_layer * num_tokens_current_layer + ].copy_(current_mask.reshape(-1)) + + cur_total = batch_size * num_tokens_current_layer + attn_metadata.spec_decoding_position_offsets[:cur_total] = new_positions.reshape(-1) + self._apply_spec_metadata(attn_metadata, batch_size, num_tokens_current_layer) + + if self._needs_mask_repack: + self._repack_mask_padded_to_packed(packed_mask, batch_size, num_tokens_current_layer) + + def update_hidden_states( + self, + cur_draft_idx, + batch_size, + step0_hs=None, + hidden_states_to_save=None, + selected_parents=None, + ): + """Manage growing-context hidden states for the MTP draft loop.""" + if cur_draft_idx == 0: + hs_dim = step0_hs.shape[-1] + self._hs_dim = hs_dim + if self._hs_write_buffer is None or self._hs_write_buffer.shape[2] != hs_dim: + self._hs_write_buffer = torch.zeros( + self._max_batch_size, + self.max_draft_len * self.K, + hs_dim, + device=step0_hs.device, + dtype=step0_hs.dtype, + ) + if self._accumulated_hs is None or self._accumulated_hs.shape[2] != hs_dim: + self._accumulated_hs = torch.zeros( + self._max_batch_size, + self.max_draft_len * self.K, + hs_dim, + device=step0_hs.device, + dtype=step0_hs.dtype, + ) + # All K depth-0 tokens share step0_hs (the parent hidden state). + self._accumulated_hs[:batch_size, : self.K] = step0_hs.unsqueeze(1).expand( + -1, self.K, -1 + ) + self._step0_hs = step0_hs + else: + num_tokens_per_req = cur_draft_idx * self.K + hs_to_save_reshaped = hidden_states_to_save.reshape(batch_size, num_tokens_per_req, -1) + self._hs_write_buffer[:batch_size, :num_tokens_per_req] = hs_to_save_reshaped + parent_offset = (cur_draft_idx - 1) * self.K + self._hs_read_map[ + :batch_size, cur_draft_idx * self.K : (cur_draft_idx + 1) * self.K + ] = parent_offset + selected_parents + num_tokens_next = (cur_draft_idx + 1) * self.K + read_idx = self._hs_read_map[:batch_size, self.K : num_tokens_next] + hs_dim = self._hs_write_buffer.shape[2] + self._accumulated_hs[:batch_size, self.K : num_tokens_next] = torch.gather( + self._hs_write_buffer[:batch_size], 1, read_idx.unsqueeze(-1).expand(-1, -1, hs_dim) + ) + + # ------------------------------------------------------------------ # + # Verification (greedy only) # + # ------------------------------------------------------------------ # + @nvtx_range("mtp_dyn.sample_and_accept_draft_tokens") + def sample_and_accept_draft_tokens(self, input_ids, logits, spec_metadata, attn_metadata): + """Greedy verification of the previous dynamic tree.""" + batch_size = attn_metadata.num_seqs + num_contexts = attn_metadata.num_contexts + num_gens = batch_size - num_contexts + N = self.tokens_per_gen_step + max_path_len = self._max_path_len + + if logits.dim() == 1: + logits = logits.unsqueeze(0) + + # Reset output buffers. + self._accepted_tokens_buf[:batch_size].zero_() + accepted_tokens = self._accepted_tokens_buf[:batch_size, :max_path_len] + self._num_accepted_tokens_buf[:batch_size].fill_(1) + num_accepted_tokens = self._num_accepted_tokens_buf[:batch_size] + self._accepted_draft_indices_tensor[:batch_size].fill_(-1) + + num_flat_tokens = logits.shape[0] + torch.argmax(logits, dim=-1, out=self._target_tokens_buf[:num_flat_tokens]) + target_tokens = self._target_tokens_buf[:num_flat_tokens] + + # Context requests: accept the sampled golden token only. + accepted_tokens[:num_contexts, 0].copy_(target_tokens[:num_contexts]) + + if num_gens > 0: + spec_tree_manager = self.spec_tree_manager + target_predict = self._target_predict_buf[:num_gens] + target_predict.copy_(target_tokens[num_contexts:].reshape(num_gens, N)) + + # No prior tree exists on bootstrap/warmup; accept the golden token. + if spec_tree_manager is None: + num_accepted_tokens[num_contexts:batch_size] = 1 + accepted_tokens[num_contexts:batch_size, 0] = target_predict[:, 0] + self._accepted_draft_indices_tensor[num_contexts:batch_size] = -1 + return accepted_tokens, num_accepted_tokens + + # candidates[:, 0] = golden token, candidates[:, 1:] = draft tokens. + candidates = self._candidates_buf[:num_gens] + candidates[:, 1:] = spec_metadata.draft_tokens.reshape(num_gens, N - 1) + candidates[:, 0] = target_predict[:, 0] + + slot_storage = spec_tree_manager.slot_storage + gen_slot_ids = slot_storage.all_ids_buf[num_contexts : num_contexts + num_gens] + tree_valid = slot_storage.has_tree[gen_slot_ids] + retrieve_packed = slot_storage.pack_retrieve_from_slots(gen_slot_ids, num_gens) + + accept_index, accept_token_num, accept_token = ( + self.tree_ops_converter.verify_dynamic_tree_greedy_out_packed( + candidates, + retrieve_packed, + target_predict, + num_gens, + self._max_path_len, + tree_valid=tree_valid, + ) + ) + tree_valid_i = tree_valid[:num_gens] + accepted_draft_count = torch.where( + tree_valid_i, + accept_token_num[:num_gens], + torch.zeros_like(accept_token_num[:num_gens]), + ) + num_accepted_tokens[num_contexts:batch_size] = (accepted_draft_count + 1).to( + torch.int32 + ) + + gen_accepted_tokens = accept_token[:num_gens].to(torch.int32) + bootstrap_accepted_tokens = torch.zeros_like(gen_accepted_tokens) + bootstrap_accepted_tokens[:, 0] = target_predict[:, 0] + accepted_tokens[num_contexts:batch_size] = torch.where( + tree_valid_i.unsqueeze(1), gen_accepted_tokens, bootstrap_accepted_tokens + ) + # Convert root/padding index 0 to draft-node sentinel -1. + gen_accepted_indices = (accept_index[:num_gens, 1:max_path_len] - 1).to(torch.int32) + self._accepted_draft_indices_tensor[num_contexts:batch_size] = torch.where( + tree_valid_i.unsqueeze(1), + gen_accepted_indices, + torch.full_like(gen_accepted_indices, -1), + ).to(torch.int32) + + num_accepted_tokens = self._apply_force_accepted_tokens( + num_accepted_tokens, num_contexts, self.max_draft_len + ) + return accepted_tokens, num_accepted_tokens + + def _accepted_leaf_intermediate_positions(self, num_accepted_tokens, num_contexts, num_gens): + """Return each accepted leaf's position in the Mamba state buffer.""" + accepted = num_accepted_tokens[num_contexts : num_contexts + num_gens].to(torch.int64) + # Column of the deepest accepted draft node, clamped to >=0 for the + # golden-only case (its value is ignored via the mask below). + draft_idx = self._accepted_draft_indices_tensor[num_contexts : num_contexts + num_gens].to( + torch.int64 + ) + last_col = (accepted - 2).clamp_(min=0, max=draft_idx.shape[1] - 1) + leaf = torch.gather(draft_idx, 1, last_col.unsqueeze(1)).squeeze(1) + 1 + # Golden-only requests (num_accepted == 1) take the root at position 0. + return torch.where(accepted > 1, leaf, torch.zeros_like(leaf)) + + @nvtx_range("mtp_dyn._relocate_kv_eagerly") + def _relocate_kv_eagerly(self, attn_metadata, batch_size): + """Move accepted draft KV from tree positions to the linear prefix.""" + cache_mgr = getattr(attn_metadata, "kv_cache_manager", None) + if cache_mgr is None or self._kv_head_dim_bytes is None: + return + if not hasattr(cache_mgr, "num_kv_heads_per_layer"): + return + + # Mamba layers have zero KV heads; relocate attention-layer KV only. + kv_heads = cache_mgr.num_kv_heads_per_layer + attn_heads = set(h for h in kv_heads if h > 0) + assert len(attn_heads) == 1, ( + "update_kv_cache_draft_token_location_2d requires uniform " + f"num_kv_heads across attention layers, got {list(kv_heads)}" + ) + attn_num_heads = attn_heads.pop() + attn_layer_offsets = [i for i, h in enumerate(kv_heads) if h > 0] + attn_num_layers = len(attn_layer_offsets) + + # Resolve the attention KV pool used by the relocation op. + pool_mapping = getattr(cache_mgr, "kv_cache_pool_mapping", None) + if pool_mapping is not None: + attn_pool_indices = set(int(pool_mapping[off][0]) for off in attn_layer_offsets) + assert len(attn_pool_indices) == 1, ( + "update_kv_cache_draft_token_location_2d requires all attention " + f"layers in one KV pool, got pools {sorted(attn_pool_indices)}" + ) + attn_pool_idx = attn_pool_indices.pop() + else: + attn_pool_idx = 0 + + pool_pointers = cache_mgr.kv_cache_pool_pointers[attn_pool_idx] + block_offsets = attn_metadata.kv_cache_block_offsets[attn_pool_idx] + + torch.ops.tensorrt_llm.update_kv_cache_draft_token_location_2d( + self._accepted_draft_indices_tensor[:batch_size], + self._num_accepted_tokens_buf[:batch_size], + attn_metadata.kv_lens_cuda[:batch_size], + True, + attn_num_layers, + attn_num_heads, + self._kv_head_dim_bytes, + cache_mgr.max_total_draft_tokens, + cache_mgr.max_attention_window_vec[0], + pool_pointers, + block_offsets, + cache_mgr.max_blocks_per_seq, + cache_mgr.tokens_per_block, + None, + ) + + # ------------------------------------------------------------------ # + # Top-level forward # + # ------------------------------------------------------------------ # + @nvtx_range("mtp_dyn.forward") + def forward( + self, + input_ids, + position_ids, + hidden_states, + logits, + attn_metadata, + spec_metadata, + draft_model, + resource_manager=None, + ): + """Run verify, cache promotion, and next-tree drafting.""" + if resource_manager is not None: + self._ensure_spec_tree_manager(resource_manager) + + batch_size = attn_metadata.num_seqs + num_contexts = attn_metadata.num_contexts + num_gens = batch_size - num_contexts + raw_logits = logits + + self._execute_guided_decoder_if_present(logits) + + # (a) Verify previous tree (greedy). Also relocates accepted KV. + accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( + input_ids, logits, spec_metadata, attn_metadata + ) + if num_gens > 0: + self._relocate_kv_eagerly(attn_metadata, batch_size) + + # Dynamic-tree Mamba states are stored by tree-node position. + if self._is_mamba_hybrid_cache is None: + self._is_mamba_hybrid_cache = isinstance( + attn_metadata.kv_cache_manager, MambaHybridCacheManager + ) + if num_gens > 0 and self._is_mamba_hybrid_cache: + accepted_leaf_positions = self._accepted_leaf_intermediate_positions( + num_accepted_tokens, num_contexts, num_gens + ) + attn_metadata.kv_cache_manager.update_mamba_states( + attn_metadata=attn_metadata, + num_accepted_tokens=num_accepted_tokens, + state_indices=attn_metadata.mamba_metadata.state_indices, + accepted_leaf_positions=accepted_leaf_positions, + ) + + # Save attn/spec metadata before the draft loop mutates it. + original_all_rank_num_tokens = attn_metadata.all_rank_num_tokens + original_force_prepare_spec_dec_tree_mask = attn_metadata.force_prepare_spec_dec_tree_mask + self._prepare_attn_metadata_for_spec_dec(attn_metadata) + attn_metadata.force_prepare_spec_dec_tree_mask = True + + # (c) Run the MTP draft tree loop -> build + store the next tree. + draft_kv_cache_manager = self.get_draft_kv_cache_manager(resource_manager) + next_draft_tokens = self._forward_draft_loop( + input_ids=input_ids, + position_ids=position_ids, + hidden_states=hidden_states, + accepted_tokens=accepted_tokens, + num_accepted_tokens=num_accepted_tokens, + attn_metadata=attn_metadata, + spec_metadata=spec_metadata, + draft_model=draft_model, + draft_kv_cache_manager=draft_kv_cache_manager, + num_contexts=num_contexts, + num_gens=num_gens, + batch_size=batch_size, + ) + + # Restore attn metadata to support cuda graph. + self._restore_attn_metadata_from_spec_dec(attn_metadata) + attn_metadata.all_rank_num_tokens = original_all_rank_num_tokens + attn_metadata.force_prepare_spec_dec_tree_mask = original_force_prepare_spec_dec_tree_mask + attn_metadata.use_spec_decoding = True + + # (d) Prepare next_new_tokens for overlap scheduler. + next_new_tokens = self._prepare_next_new_tokens( + accepted_tokens, + next_draft_tokens, + spec_metadata.batch_indices_cuda, + batch_size, + num_accepted_tokens, + ) + + return { + "logits": raw_logits, + "new_tokens": accepted_tokens, + "new_tokens_lens": num_accepted_tokens, + "next_draft_tokens": next_draft_tokens, + "next_new_tokens": next_new_tokens, + "accepted_draft_tokens_indices": self._accepted_draft_indices_tensor[:batch_size], + } + + # ------------------------------------------------------------------ # + # Step-0 drafter-input repack (dynamic tree) # + # ------------------------------------------------------------------ # + @nvtx_range("mtp_dyn._prepare_step0_drafter_inputs") + def _prepare_step0_drafter_inputs( + self, + input_ids, + position_ids, + last_tokens_idx, + hidden_states, + accepted_tokens, + attn_metadata, + ): + """Repack step-0 drafter inputs to accepted-path layout.""" + num_contexts = attn_metadata.num_contexts + batch_size = attn_metadata.num_seqs + num_gens = batch_size - num_contexts + num_ctx_tokens = attn_metadata.num_ctx_tokens + + # Match MTPEagleWorker context input repack. + input_ids_ctx = self._prepare_context_input_ids( + input_ids, num_ctx_tokens, last_tokens_idx, accepted_tokens, num_contexts + ) + + if num_gens > 0: + max_path_len = self._max_path_len + num_gen_tokens = num_gens * max_path_len + + hidden_dim = hidden_states.shape[-1] + if ( + self._step0_hidden_states_buf is None + or self._step0_hidden_states_buf.shape[-1] != hidden_dim + ): + self._step0_hidden_states_buf = torch.zeros( + self._step0_input_ids_buf.shape[0], + hidden_dim, + dtype=hidden_states.dtype, + device="cuda", + ) + + # Accepted path includes the golden token at column 0. + accept_token = accepted_tokens[num_contexts:batch_size] + + BLOCK_H = triton.next_power_of_2(hidden_dim) + _gather_repack_step0_kernel[(num_gens * max_path_len,)]( + hidden_states, + accept_token, + position_ids, + self._accepted_draft_indices_tensor[num_contexts:batch_size], + self._num_accepted_tokens_buf, + self._step0_hidden_states_buf, + self._step0_input_ids_buf, + self._step0_position_ids_buf, + self._gather_ids_buf, + num_ctx_tokens, + num_contexts, + self.tokens_per_gen_step, + max_path_len, + self.max_draft_len, + hidden_dim, + num_ctx_tokens, # gather_id references combined [ctx|gen] tensor + BLOCK_H=BLOCK_H, + ) + + input_ids = torch.cat( + [input_ids_ctx, self._step0_input_ids_buf[:num_gen_tokens]], dim=0 + ) + position_ids = torch.cat( + [position_ids[:num_ctx_tokens], self._step0_position_ids_buf[:num_gen_tokens]], + dim=0, + ) + hidden_states = torch.cat( + [hidden_states[:num_ctx_tokens], self._step0_hidden_states_buf[:num_gen_tokens]], + dim=0, + ) + + attn_metadata._seq_lens[num_contexts:batch_size].fill_(max_path_len) + attn_metadata._seq_lens_cuda[num_contexts:batch_size].fill_(max_path_len) + attn_metadata.on_update() + else: + # Context-only (warmup): no gen tokens to repack. + input_ids = input_ids_ctx + + return { + "input_ids": input_ids, + "position_ids": position_ids, + "hidden_states": hidden_states, + "attn_metadata": attn_metadata, + } + + # ------------------------------------------------------------------ # + # MTP draft tree loop # + # ------------------------------------------------------------------ # + def _forward_draft_loop( + self, + input_ids, + position_ids, + hidden_states, + accepted_tokens, + num_accepted_tokens, + attn_metadata, + spec_metadata, + draft_model, + draft_kv_cache_manager, + num_contexts, + num_gens, + batch_size, + ): + """Draft the next dynamic tree with growing context.""" + spec_tree_manager = self.spec_tree_manager + + assert batch_size <= self._max_batch_size, ( + f"batch_size {batch_size} exceeds pre-allocated max_batch_size {self._max_batch_size}" + ) + + # Step 0: run MTP over accepted-path rows. + position_ids, last_tokens_idx = self.prepare_position_ids_and_last_tokens( + position_ids, attn_metadata.seq_lens_cuda + ) + inputs = self._prepare_step0_drafter_inputs( + input_ids=input_ids, + position_ids=position_ids, + last_tokens_idx=last_tokens_idx, + hidden_states=hidden_states, + accepted_tokens=accepted_tokens, + attn_metadata=attn_metadata, + ) + + # Reset verify-time tree metadata to accepted-path width. + num_step0_tokens = self._max_path_len + if attn_metadata.spec_decoding_generation_lengths is not None: + total = num_gens * num_step0_tokens + dst = attn_metadata.spec_decoding_position_offsets[:total].view( + num_gens, num_step0_tokens + ) + dst.copy_(self._causal_offs[:num_step0_tokens].unsqueeze(0).expand(num_gens, -1)) + self._apply_spec_metadata(attn_metadata, num_gens, num_step0_tokens) + packed_mask = attn_metadata.spec_decoding_packed_mask + packed_mask[:num_gens].zero_() + packed_mask[:num_gens, :num_step0_tokens, 0] = self._step0_causal_mask[ + :num_step0_tokens + ] + if self._needs_mask_repack: + self._repack_mask_padded_to_packed(packed_mask, num_gens, num_step0_tokens) + attn_metadata.use_spec_decoding = num_gens > 0 + if num_gens > 0 and hasattr(attn_metadata, "kv_lens_cuda"): + attn_metadata.kv_lens_cuda[num_contexts:batch_size] -= self._kv_correction + self._refresh_blackwell_tree_mask_metadata(attn_metadata) + if spec_metadata.all_rank_num_tokens is not None: + # Keep attention/MoE token counts aligned with step-0 repack. + attn_metadata.all_rank_num_tokens = spec_metadata.all_rank_num_tokens + + with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): + hidden_states = draft_model.mtp_layers[0]( + embed_tokens=draft_model.embed_tokens, + all_rank_num_tokens=spec_metadata.all_rank_num_tokens, + **inputs, + ) + + # Gather each request's root hidden state for depth-0 expansion. + self._gather_ids_buf[:num_contexts].copy_(last_tokens_idx[:num_contexts]) + gather_ids = self._gather_ids_buf[:batch_size] + + step0_hs = hidden_states[gather_ids] + logits = draft_model.mtp_layers[0].shared_head( + step0_hs, draft_model.lm_head, attn_metadata, True + ) + + new_draft_tokens, new_draft_scores = self.sample( + logits, self.K, draft_model=draft_model + ) + previous_draft_scores = self.update_draft_tokens_and_scores( + cur_draft_idx=0, + new_draft_tokens=new_draft_tokens, + new_draft_scores=new_draft_scores, + previous_draft_scores=None, + batch_size=batch_size, + attn_metadata=attn_metadata, + ) + self.update_hidden_states(cur_draft_idx=0, batch_size=batch_size, step0_hs=step0_hs) + self._prepare_draft_layer_metadata( + 0, + attn_metadata, + batch_size, + gather_ids, + num_contexts, + num_gens, + num_accepted_tokens, + inputs, + ) + + # Subsequent layers grow the tree. + for layer_idx in range(1, self.max_draft_len): + num_tokens_per_req = layer_idx * self.K + num_infer_tokens = batch_size * num_tokens_per_req + subseq_all_rank_num_tokens = None + if spec_metadata.all_rank_num_seqs is not None: + # Token counts scale with the current tree width. + subseq_all_rank_num_tokens = [ + n * num_tokens_per_req for n in spec_metadata.all_rank_num_seqs + ] + attn_metadata.all_rank_num_tokens = subseq_all_rank_num_tokens + + inp_hs = self._accumulated_hs[:batch_size, :num_tokens_per_req, :].reshape( + num_infer_tokens, -1 + ) + inp_ids = self.draft_tokens_buffer[:batch_size, :num_tokens_per_req].reshape(-1) + inp_pos = self.position_ids_buffer[:batch_size, :num_tokens_per_req].reshape(-1) + layer_inputs = { + "input_ids": inp_ids, + "position_ids": inp_pos, + "hidden_states": inp_hs, + "attn_metadata": attn_metadata, + } + hidden_states = draft_model.mtp_layers[0]( + embed_tokens=draft_model.embed_tokens, + all_rank_num_tokens=subseq_all_rank_num_tokens + or spec_metadata.subseq_all_rank_num_tokens, + **layer_inputs, + ) + + # Take the last K hidden states per request (the new leaves). + hs_reshaped = hidden_states.reshape(batch_size, num_tokens_per_req, -1) + selected_hs = hs_reshaped[:, -self.K :, :].reshape(batch_size * self.K, -1) + logits = draft_model.mtp_layers[0].shared_head( + selected_hs, draft_model.lm_head, attn_metadata, True + ) + + new_draft_tokens, new_draft_scores = self.sample( + logits, self.K, draft_model=draft_model + ) + new_draft_tokens = new_draft_tokens.reshape(batch_size, self.K, self.K) + new_draft_scores = new_draft_scores.reshape(batch_size, self.K, self.K) + + previous_draft_scores = self.update_draft_tokens_and_scores( + cur_draft_idx=layer_idx, + new_draft_tokens=new_draft_tokens, + new_draft_scores=new_draft_scores, + previous_draft_scores=previous_draft_scores, + batch_size=batch_size, + attn_metadata=attn_metadata, + ) + self.update_hidden_states( + cur_draft_idx=layer_idx, + batch_size=batch_size, + hidden_states_to_save=hidden_states, + selected_parents=self._last_selected_parents, + ) + self._prepare_draft_layer_metadata(layer_idx, attn_metadata, batch_size) + + # Resample the final tree and build it into slot_storage. + real_draft_tokens, topk_score_indices = self.resampling_final_draft_tokens(batch_size) + + if spec_tree_manager is not None and num_gens > 0: + self.tree_ops_converter.build_dynamic_tree( + history_draft_tokens_parent_buffer=self.history_draft_tokens_parent_buffer[ + num_contexts:batch_size + ], + topk_score_indices=topk_score_indices[num_contexts:], + tree_mask=spec_tree_manager.spec_dec_packed_mask[:num_gens], + positions=spec_tree_manager.spec_dec_position_offsets[:num_gens], + retrieve_index=spec_tree_manager.retrieve_index[:num_gens], + retrieve_next_token=spec_tree_manager.retrieve_next_token[:num_gens], + retrieve_next_sibling=spec_tree_manager.retrieve_next_sibling[:num_gens], + use_packed_mask=True, + ) + slot_storage = spec_tree_manager.slot_storage + gen_slots = slot_storage.all_ids_buf[num_contexts:batch_size] + spec_tree_manager.scatter_to_slot_storage(slot_storage, gen_slots, num_gens) + + return real_draft_tokens + + def _prepare_draft_layer_metadata( + self, + cur_draft_idx, + attn_metadata, + batch_size, + gather_ids=None, + num_contexts=0, + num_gens=0, + num_accepted_tokens=None, + inputs=None, + ): + """Set attn_metadata seq_lens/kv_lens for the next draft layer.""" + if cur_draft_idx == 0: + base_pos = inputs["position_ids"][gather_ids] + 1 + self.position_ids_buffer[:batch_size, : self.K] = base_pos.unsqueeze(1).expand( + -1, self.K + ) + + attn_metadata._seq_lens[:batch_size].fill_(self.K) + attn_metadata._seq_lens_cuda[:batch_size].fill_(self.K) + attn_metadata.on_update() + + if inputs["attn_metadata"].kv_cache_manager is not None: + attn_metadata.host_request_types[: attn_metadata.num_contexts].fill_(1) + attn_metadata.num_contexts = 0 + + if hasattr(attn_metadata, "kv_lens_cuda"): + # Rewind only unaccepted verify tokens; draft KV is added later. + if num_gens > 0: + attn_metadata.kv_lens_cuda[num_contexts:batch_size] -= ( + self._max_path_len + ) - num_accepted_tokens[num_contexts:batch_size] + attn_metadata.kv_lens_cuda[:batch_size] += self.K + attn_metadata.use_spec_decoding = True + self._refresh_blackwell_tree_mask_metadata(attn_metadata) + else: + num_tokens_previous_layer = cur_draft_idx * self.K + num_tokens_current_layer = self.K * (cur_draft_idx + 1) + prev_pos = self.position_ids_buffer[:batch_size, :num_tokens_previous_layer] + self.position_ids_buffer[ + :batch_size, num_tokens_previous_layer:num_tokens_current_layer + ] = prev_pos[:, -self.K :] + 1 + attn_metadata._seq_lens[:batch_size].fill_(num_tokens_current_layer) + attn_metadata._seq_lens_cuda[:batch_size].fill_(num_tokens_current_layer) + attn_metadata.on_update() + if hasattr(attn_metadata, "kv_lens_cuda"): + attn_metadata.kv_lens_cuda[:batch_size] += self.K + self._refresh_blackwell_tree_mask_metadata(attn_metadata) + + +class MTPEagleDynamicTreeResourceManager(BaseResourceManager): + """Resource manager for MTP dynamic-tree mode.""" + + hidden_states: Optional[torch.Tensor] = None + + def __init__( + self, + config: "MTPDecodingConfig", + dtype: torch.dtype, + hidden_size: int, + max_num_requests: int, + sa_manager=None, + ): + from .spec_tree_manager import SpecTreeManager + + self.max_num_requests = max_num_requests + self.spec_tree_manager = SpecTreeManager( + max_num_requests=max_num_requests, + use_dynamic_tree=True, + max_draft_len=config.max_draft_len, + max_total_draft_tokens=config.tokens_per_gen_step - 1, + eagle_choices=None, + dynamic_tree_max_topK=config.dynamic_tree_max_topK, + ) + # MTP hidden-state slot pools (needed by MTPEagleWorker drafter inputs). + self._mtp_hidden_states_manager = MTPHiddenStatesManager( + config, dtype, hidden_size, max_num_requests, sa_manager=sa_manager + ) + + # Expose the MTPHiddenStatesManager surface MTPSpecMetadata expects. + @property + def slot_manager(self): + return self._mtp_hidden_states_manager.slot_manager + + @property + def mtp_past_hidden_states_pool(self): + return self._mtp_hidden_states_manager.mtp_past_hidden_states_pool + + @property + def mtp_past_tokens_pool(self): + return self._mtp_hidden_states_manager.mtp_past_tokens_pool + + @property + def sa_manager(self): + return self._mtp_hidden_states_manager.sa_manager + + def prepare_resources(self, scheduled_batch: ScheduledRequests): + self._mtp_hidden_states_manager.prepare_resources(scheduled_batch) + + def update_resources(self, scheduled_batch: ScheduledRequests): + self._mtp_hidden_states_manager.update_resources(scheduled_batch) + + def free_resources(self, request: LlmRequest): + # Clear tree validity for the freed slot, then free the MTP slot. + if request.py_seq_slot is not None: + self.spec_tree_manager.slot_storage.mark_invalid(request.py_seq_slot) + self._mtp_hidden_states_manager.free_resources(request) + + def add_dummy_requests(self, request_ids: List[int]): + # Dummies still need MTP hidden-state slots. + self._mtp_hidden_states_manager.add_dummy_requests(request_ids) + + def shutdown(self): + self._mtp_hidden_states_manager.shutdown() + + def get_max_resource_count(self) -> int: + return self.max_num_requests + + def get_needed_resource_to_completion(self, request: LlmRequest): + return 0 diff --git a/tensorrt_llm/_torch/speculative/spec_tree_manager.py b/tensorrt_llm/_torch/speculative/spec_tree_manager.py index 74dd622826f1..545b5bc7bb06 100644 --- a/tensorrt_llm/_torch/speculative/spec_tree_manager.py +++ b/tensorrt_llm/_torch/speculative/spec_tree_manager.py @@ -16,24 +16,40 @@ class DynamicTreeSlotStorage: Buffers are [S, ...] where S = num_slots + 1 (+1 for CUDA graph dummy). """ - def __init__(self, num_slots: int, n_dt: int, mask_width: int): + def __init__(self, + num_slots: int, + n_dt: int, + mask_width: int, + top_k: int = 1): S = num_slots + 1 self.dummy_slot_id = num_slots - # Slot buffers — C++ kernel writes directly via slotIds - self.packed_mask = torch.zeros((S, n_dt, mask_width), - dtype=torch.int32, - device='cuda') - self.position_offsets = torch.zeros((S, n_dt), - dtype=torch.int32, - device='cuda') + # Bootstrap/reused slots may not have a tree yet; keep their metadata + # as a valid linear chain so verification kernels can read it directly. + no_tree_position_offsets, no_tree_packed_mask = self._make_kary_tree_metadata( + n_dt, mask_width, top_k=1) + self.position_offsets = no_tree_position_offsets.unsqueeze(0).repeat( + S, 1).contiguous() + self.packed_mask = no_tree_packed_mask.unsqueeze(0).repeat( + S, 1, 1).contiguous() + self._no_tree_position_offsets = no_tree_position_offsets + self._no_tree_packed_mask = no_tree_packed_mask + + # CUDA-graph dummies use a deterministic K-ary tree, matching real + # dynamic-tree mask/position shapes without depending on request state. + dummy_position_offsets, dummy_packed_mask = self._make_kary_tree_metadata( + n_dt, mask_width, top_k) + self.position_offsets[self.dummy_slot_id] = dummy_position_offsets + self.packed_mask[self.dummy_slot_id] = dummy_packed_mask self.retrieve_index = torch.zeros((S, n_dt), dtype=torch.int32, device='cuda') - self.retrieve_next_token = torch.full((S, n_dt), - -1, - dtype=torch.int32, - device='cuda') + + # Mamba verify reads next links unconditionally, so no-tree rows must be + # valid linear chains instead of sentinels. + self._no_tree_next_token = self._make_no_tree_next_token(n_dt) + self.retrieve_next_token = self._no_tree_next_token.unsqueeze(0).repeat( + S, 1) self.retrieve_next_sibling = torch.full((S, n_dt), -1, dtype=torch.int32, @@ -57,6 +73,43 @@ def __init__(self, num_slots: int, n_dt: int, mask_width: int): dtype=torch.int32, device='cuda') + @staticmethod + def _make_kary_tree_metadata( + n_dt: int, mask_width: int, + top_k: int) -> tuple[torch.Tensor, torch.Tensor]: + top_k = max(int(top_k), 1) + token_ids = torch.arange(n_dt, device='cuda') + parents = torch.where(token_ids > 0, (token_ids - 1) // top_k, + token_ids) + ancestor_chain = torch.empty((n_dt, n_dt), + dtype=torch.long, + device='cuda') + current = token_ids + for depth in range(n_dt): + ancestor_chain[:, depth] = current + current = parents[current] + + # Pack bits directly from the parent chain instead of materializing a + # dense bool mask and repacking it. + valid_ancestors = torch.ones((n_dt, n_dt), + dtype=torch.bool, + device='cuda') + valid_ancestors[:, 1:] = ancestor_chain[:, 1:] != ancestor_chain[:, :-1] + bit_values = (1 << (ancestor_chain % 32)).to(torch.int32) + bit_values.masked_fill_(~valid_ancestors, 0) + packed_mask = torch.zeros((n_dt, mask_width), + dtype=torch.int32, + device='cuda') + packed_mask.scatter_add_(1, ancestor_chain // 32, bit_values) + position_offsets = valid_ancestors.sum(-1).to(torch.int32) - 1 + return position_offsets, packed_mask + + @staticmethod + def _make_no_tree_next_token(n_dt: int) -> torch.Tensor: + next_token = torch.arange(1, n_dt + 1, dtype=torch.int32, device='cuda') + next_token[n_dt - 1] = -1 + return next_token + def fill_all_slot_ids(self, context_requests, generation_requests): """Fill all_ids_buf for full batch [ctx | gen] via one HtoD copy.""" dummy_slot = self.dummy_slot_id @@ -81,12 +134,12 @@ def mark_valid(self, slot_ids, count): self.has_tree.narrow(0, self.dummy_slot_id, 1).fill_(False) def mark_invalid(self, slot_id): - """Clear validity and reset slot data.""" + """Clear validity and restore valid no-tree metadata.""" self.has_tree[slot_id] = False - self.packed_mask[slot_id] = 0 - self.position_offsets[slot_id] = 0 + self.packed_mask[slot_id] = self._no_tree_packed_mask + self.position_offsets[slot_id] = self._no_tree_position_offsets self.retrieve_index[slot_id] = 0 - self.retrieve_next_token[slot_id] = -1 + self.retrieve_next_token[slot_id] = self._no_tree_next_token self.retrieve_next_sibling[slot_id] = -1 def pack_retrieve_from_slots(self, slot_ids, count): @@ -284,6 +337,7 @@ def init_tree_info_for_dynamic_tree(self): num_slots=self.num_trees, n_dt=num_draft_with_root, mask_width=mask_width, + top_k=self.dynamic_tree_max_topK, ) def scatter_to_slot_storage(self, ss, gen_slots, num_gens): diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 4f9c5af8846a..d800751b209f 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -27,6 +27,8 @@ from .eagle3_dynamic_tree import Eagle3OneModelDynamicTreeWorker from .model_drafter import ModelDrafter from .mtp import MTPHiddenStatesManager, MTPSampler, MTPSpecMetadata, MTPWorker +from .mtp_dynamic_tree import (MTPEagleDynamicTreeResourceManager, + MTPEagleDynamicTreeWorker) from .ngram import NGramDrafter, NGramPoolManager from .pard import PARDSpecMetadata, PARDWorker from .sa_worker import SASampler, SASpecMetadata, SAWorker @@ -112,6 +114,7 @@ def get_spec_metadata(spec_config, num_seq_slots=num_seq_slots, draft_vocab_size=draft_vocab_size, spec_resource_manager=spec_resource_manager, + use_dynamic_tree=getattr(spec_config, 'use_dynamic_tree', False), ) if spec_config.spec_dec_mode.is_mtp_vanilla(): return MTPSpecMetadata( @@ -273,6 +276,15 @@ def get_spec_resource_manager(model_engine, draft_model_engine=None): if sa_cfg is not None: sa_manager = SuffixAutomatonManager(sa_cfg, max_num_requests, max_seq_len) + # Dynamic tree combines SpecTreeManager with MTP hidden-state slots. + if getattr(spec_config, 'use_dynamic_tree', False): + return MTPEagleDynamicTreeResourceManager( + spec_config, + model_config.torch_dtype, + model_config.hidden_size, + max_num_requests, + sa_manager=sa_manager, + ) if spec_config.use_relaxed_acceptance_for_thinking or sa_manager is not None: # Unified resource manager: the unified worker reads # ``relaxed_delta_pool`` from ``Eagle3ResourceManager`` (mirrors the @@ -361,7 +373,10 @@ def get_spec_decoder( # MTP Eagle one-model now uses the same sampler as Eagle3 one-model. return Eagle3OneModelSampler(sampler_args, spec_config=spec_config) if spec_config.spec_dec_mode.is_mtp_vanilla(): - return MTPSampler(sampler_args, nextn=spec_config.max_draft_len) + nextn = spec_config.max_draft_len + if getattr(spec_config, "use_dynamic_tree", False): + nextn = spec_config.max_total_draft_tokens + return MTPSampler(sampler_args, nextn=nextn) if spec_config.spec_dec_mode.is_eagle3( ) or spec_config.spec_dec_mode.is_mtp_eagle(): # TorchSampler handles Eagle3 gracefully, by integrating d2t into the sampling process @@ -449,6 +464,11 @@ def get_spec_worker(spec_config, use_separate_draft_kv_cache, mapping=mapping) if spec_dec_mode.is_mtp_eagle_one_model(): + if getattr(spec_config, 'use_dynamic_tree', False): + return MTPEagleDynamicTreeWorker(spec_config, + model_config, + use_separate_draft_kv_cache, + mapping=mapping) return MTPEagleWorker(spec_config, model_config, use_separate_draft_kv_cache, @@ -536,7 +556,8 @@ def update_spec_config_from_model_config(spec_config, model_config): f"using max_draft_len={effective_draft_len} draft tokens.") spec_config.max_draft_len = effective_draft_len - spec_config.max_total_draft_tokens = spec_config.max_draft_len + if not spec_config.use_dynamic_tree: + spec_config.max_total_draft_tokens = spec_config.max_draft_len def update_spec_config_from_loaded_model(spec_config, model) -> None: diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 57c3de25dcfe..4ec8383b5a51 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1971,8 +1971,10 @@ class EagleDecodingConfig(DecodingBaseConfig): ) dynamic_tree_max_topK: Optional[int] = Field( default=None, - description="The topK value for each layer when dynamic tree is enabled." - ) + description= + "The topK value for each layer when dynamic tree is enabled. Required " + "when use_dynamic_tree is True; ignored (with a warning) when " + "use_dynamic_tree is False.") num_eagle_layers: Optional[int] = Field( default=None, description= @@ -2048,9 +2050,17 @@ def validate_eagle_config(self) -> 'EagleDecodingConfig': # So the number of choices also represents the number of max draft nodes. self.max_total_draft_tokens = len(self.eagle_choices) + # Dynamic tree is enabled only by an explicit use_dynamic_tree=True; + # dynamic_tree_max_topK alone does not turn it on. + if not self.use_dynamic_tree and self.dynamic_tree_max_topK is not None: + logger.warning( + "dynamic_tree_max_topK is set but use_dynamic_tree is False; " + "ignoring dynamic_tree_max_topK and using the linear draft path." + ) + self.dynamic_tree_max_topK = None + # Dynamic tree logic - if self.use_dynamic_tree or self.dynamic_tree_max_topK is not None: - self.use_dynamic_tree = True + if self.use_dynamic_tree: if self.eagle_choices is not None: raise ValueError( "If use_dynamic_tree is True, eagle_choices should be None") @@ -2070,16 +2080,12 @@ def validate_eagle_config(self) -> 'EagleDecodingConfig': logger.warning( f"max_total_draft_tokens is not provided, use the default value {default_max_total_draft_tokens} (default_max_total_draft_tokens = dynamic_tree_max_topK * max_draft_len)" ) - else: - if self.max_total_draft_tokens < self.max_draft_len: - raise ValueError( - f"max_total_draft_tokens ({self.max_total_draft_tokens}) should be >= max_draft_len ({self.max_draft_len})" - ) - if self.max_total_draft_tokens > self.dynamic_tree_max_topK * self.max_draft_len: - raise ValueError( - f"max_total_draft_tokens ({self.max_total_draft_tokens}) should be <= " - f"dynamic_tree_max_topK * max_draft_len ({self.dynamic_tree_max_topK * self.max_draft_len})" - ) + elif not (self.max_draft_len <= self.max_total_draft_tokens <= + default_max_total_draft_tokens): + raise ValueError( + f"max_total_draft_tokens ({self.max_total_draft_tokens}) must be in " + f"[max_draft_len ({self.max_draft_len}), dynamic_tree_max_topK * " + f"max_draft_len ({default_max_total_draft_tokens})]") # Linear tree if self.max_total_draft_tokens is None: @@ -2425,6 +2431,22 @@ class MTPDecodingConfig(DecodingBaseConfig): "When using EAGLE-style MTP, use faster one-model implementation (drafter as submodule) vs two-model." ) + use_dynamic_tree: bool = Field( + default=False, + description= + "Enable EAGLE-style dynamic-tree drafting for one-model MTP. When True, " + "each draft step expands dynamic_tree_max_topK candidates per node and the " + "tree is verified against the target, instead of a linear chain.") + dynamic_tree_max_topK: Optional[int] = Field( + default=None, + description= + "Top-K candidates expanded per node per draft layer when use_dynamic_tree " + "is enabled. Required when use_dynamic_tree is True; ignored (with a " + "warning) when use_dynamic_tree is False.") + + # Internal max batch size for dynamic-tree worker buffers. + _max_batch_size: Optional[int] = PrivateAttr(default=None) + sa_config: Optional[SAEnhancerConfig] = Field( default=None, status="beta", @@ -2469,15 +2491,43 @@ def _remap_deprecated_num_nextn_predict_layers(cls, data): @model_validator(mode="after") def set_max_total_draft_tokens(self): - # Leave max_draft_len as None ("use the model's num_nextn_predict_layers") - # when the user doesn't set it; update_spec_config_from_model_config - # resolves it from the checkpoint before the model runs. When the user - # does set it, validate and mirror to max_total_draft_tokens (current MTP - # only supports a linear tree). + # None means update_spec_config_from_model_config resolves it from checkpoint. if self.max_draft_len is not None: if self.max_draft_len <= 0: raise ValueError("max_draft_len must be > 0 for MTP") - self.max_total_draft_tokens = self.max_draft_len + + # Dynamic tree is enabled only by an explicit use_dynamic_tree=True; + # dynamic_tree_max_topK alone does not turn it on. + if not self.use_dynamic_tree and self.dynamic_tree_max_topK is not None: + logger.warning( + "dynamic_tree_max_topK is set but use_dynamic_tree is False; " + "ignoring dynamic_tree_max_topK and using the linear draft path." + ) + self.dynamic_tree_max_topK = None + + # Dynamic tree defaults max_total_draft_tokens to topK * max_draft_len. + if self.use_dynamic_tree: + if self.max_draft_len is None: + raise ValueError( + "max_draft_len must be set when use_dynamic_tree is True") + if self.dynamic_tree_max_topK is None or self.dynamic_tree_max_topK <= 0: + raise ValueError( + "dynamic_tree_max_topK must be > 0 when use_dynamic_tree is True" + ) + default_max_total_draft_tokens = self.dynamic_tree_max_topK * self.max_draft_len + if self.max_total_draft_tokens is None: + self.max_total_draft_tokens = default_max_total_draft_tokens + logger.warning( + f"max_total_draft_tokens is not provided, use the default value {default_max_total_draft_tokens} (default_max_total_draft_tokens = dynamic_tree_max_topK * max_draft_len)" + ) + elif not (self.max_draft_len <= self.max_total_draft_tokens <= + default_max_total_draft_tokens): + raise ValueError( + f"max_total_draft_tokens ({self.max_total_draft_tokens}) must be in " + f"[max_draft_len ({self.max_draft_len}), dynamic_tree_max_topK * " + f"max_draft_len ({default_max_total_draft_tokens})]") + elif self.max_draft_len is not None: + self.max_total_draft_tokens = self.max_draft_len # linear chain return self @model_validator(mode="after") diff --git a/tests/unittest/_torch/speculative/test_eagle3.py b/tests/unittest/_torch/speculative/test_eagle3.py index 4e7faecbbea2..33cf3d635661 100644 --- a/tests/unittest/_torch/speculative/test_eagle3.py +++ b/tests/unittest/_torch/speculative/test_eagle3.py @@ -11,7 +11,8 @@ import torch from test_common.llm_data import with_mocked_hf_download_for_single_gpu from utils.llm_data import llm_models_root -from utils.util import skip_blackwell, skip_num_gpus_less_than +from utils.util import (skip_blackwell, skip_num_gpus_less_than, + skip_pre_blackwell) from tensorrt_llm import LLM, SamplingParams from tensorrt_llm._torch.attention_backend.trtllm import TrtllmAttentionMetadata @@ -23,7 +24,7 @@ from tensorrt_llm._torch.speculative.eagle3 import Eagle3OneModelSpecMetadata from tensorrt_llm.executor.request import LoRARequest from tensorrt_llm.llmapi import (CudaGraphConfig, Eagle3DecodingConfig, - KvCacheConfig) + KvCacheConfig, MoeConfig, MTPDecodingConfig) from tensorrt_llm.lora_helper import LoraConfig sys.path.append(os.path.join(os.path.dirname(__file__), '..')) @@ -1077,6 +1078,84 @@ def test_llama_eagle3_rejection_sampling_modes(use_dynamic_tree: bool, assert len(results[0].outputs[0].token_ids) > 0 +@pytest.mark.parametrize("disable_overlap_scheduler", [False, True]) +@pytest.mark.parametrize("use_cuda_graph", [False, True]) +@pytest.mark.high_cuda_memory +@skip_pre_blackwell +def test_nemotron_super_mtp_dynamic_tree_dl6_k10_dt31( + use_cuda_graph: bool, disable_overlap_scheduler: bool): + if torch.cuda.device_count() < 8: + pytest.skip("Nemotron Super dynamic-tree MTP test requires 8 GPUs") + + models_path = llm_models_root() + model_path = f"{models_path}/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4" + + max_batch_size = 1 + max_draft_len = 6 + kv_cache_config = KvCacheConfig(enable_block_reuse=False, + mamba_ssm_cache_dtype="float16", + free_gpu_memory_fraction=0.8) + cuda_graph_config = CudaGraphConfig( + batch_sizes=[1]) if use_cuda_graph else None + + llm_common_config = dict( + model=model_path, + tensor_parallel_size=8, + moe_expert_parallel_size=8, + pipeline_parallel_size=1, + moe_config=MoeConfig(backend="TRTLLM"), + disable_overlap_scheduler=disable_overlap_scheduler, + cuda_graph_config=cuda_graph_config, + max_batch_size=max_batch_size, + kv_cache_config=kv_cache_config, + max_seq_len=8192, + ) + spec_config = MTPDecodingConfig(max_draft_len=max_draft_len, + mtp_eagle_one_model=True, + use_dynamic_tree=True, + dynamic_tree_max_topK=10, + max_total_draft_tokens=31) + + llm_spec = LLM(**llm_common_config, speculative_config=spec_config) + prompt = llm_spec.tokenizer.apply_chat_template( + [{ + "role": "user", + "content": "The future of AI is" + }], + tokenize=False, + add_generation_prompt=True, + ) + tok_ids = llm_spec.tokenizer.encode(prompt) + + sampling_params = SamplingParams(max_tokens=128, temperature=0) + num_tokens = 0 + num_drafted = 0 + num_accepted = 0 + for output in llm_spec.generate_async(tok_ids, + sampling_params, + streaming=True): + new_tokens = output.outputs[0].token_ids + num_drafted += max_draft_len + num_accepted += len(new_tokens) - num_tokens - 1 + num_tokens = len(new_tokens) + + accept_rate = num_accepted / num_drafted + assert accept_rate > 0.20 + + sampling_params = SamplingParams(max_tokens=10, temperature=0) + results_spec = llm_spec.generate([prompt], sampling_params) + generated_text_spec = [result.outputs[0].text for result in results_spec] + llm_spec.shutdown() + + llm_ref = LLM(**llm_common_config) + results_ref = llm_ref.generate([prompt], sampling_params) + generated_text_ref = [result.outputs[0].text for result in results_ref] + llm_ref.shutdown() + + for text_spec, text_ref in zip(generated_text_spec, generated_text_ref): + assert text_spec == text_ref + + @pytest.mark.parametrize("use_cuda_graph", [True, False]) def test_eagle3_lora(use_cuda_graph: bool): """Test LoRA with 3 requests and max_batch_size=4.