Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
4f2ad5d
[None][refactor] isolate sparse attention implementations
lfr-0531 Apr 27, 2026
94fc988
[None][refactor] unify sparse MLA backend dispatch
lfr-0531 Jul 7, 2026
25dff25
[None][refactor] route sparse MLA through module facade
lfr-0531 Jul 7, 2026
8675a55
[None][refactor] refine sparse attention module hooks
lfr-0531 Jul 22, 2026
fbd7f5b
[None][refactor] consolidate sparse attention runtime interfaces
lfr-0531 Jul 23, 2026
fc230c1
[None][refactor] simplify shared DSA top-k storage
lfr-0531 Jul 23, 2026
e78e616
[None][test] preserve inference mode in sparse MLA tests
lfr-0531 Jul 23, 2026
5e1f9a2
[None][fix] fix sparse attention fallback paths
lfr-0531 Jul 24, 2026
7638bff
[None][fix] fix sparse imports and PP validation
lfr-0531 Jul 24, 2026
6c91b8d
[None][refactor] adopt standard linting for sparse attention
lfr-0531 Jul 24, 2026
06d27b4
[None][fix] fix sparse attention compile interfaces
lfr-0531 Jul 26, 2026
6f6eef7
[None][refactor] clarify sparse attention hook contracts
lfr-0531 Jul 27, 2026
b5d54c8
[None][docs] clarify MiniMax M3 sparse integration
lfr-0531 Jul 27, 2026
dba871c
[None][refactor] add typed sparse attention hook adapters
lfr-0531 Jul 27, 2026
6f0d72d
[None][test] update DeepSeek V4 hook assertion
lfr-0531 Jul 28, 2026
3daa05b
[None][refactor] refine sparse attention hook contracts
lfr-0531 Jul 28, 2026
d1db820
[None][fix] preserve DSA MTP top-k sharing after rebase
lfr-0531 Aug 1, 2026
a17e6f6
[None][fix] preserve sparse MLA auxiliary streams
lfr-0531 Aug 3, 2026
81c0d50
[None][fix] Pass DeepSeek-V4 backend params to indexer
lfr-0531 Aug 4, 2026
8bcdc66
[None][fix] handle phase-specific DSA MTP top-k stash
lfr-0531 Aug 4, 2026
41ef8a1
[None][fix] preserve masked DSA indexer cache allocation
lfr-0531 Aug 5, 2026
a181ebe
[None][fix] preserve DSA runtime metadata after rebase
lfr-0531 Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions cpp/tensorrt_llm/common/attentionOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1224,11 +1224,11 @@ int AttentionOp::mlaGeneration(
// Set the following parameters if sparseAttention is used.
if (useSparseMLA())
{
bool const useDynamicSparseMLA = mRuntimeSparseAttentionParams.sparse_mla_topk_lens != nullptr;
bool const useDynamicSparseMLA = mRuntimeSparseAttentionParams.sparse_attn_kv_lens != nullptr;
tllmRunnerParams.mSparseAttention
= useDynamicSparseMLA ? SparseType::DynamicTokenSparse : SparseType::StaticTokenSparse;
tllmRunnerParams.mSparseTopK = mRuntimeSparseAttentionParams.num_sparse_topk;
tllmRunnerParams.ptrSparseMlaTopKLens = mRuntimeSparseAttentionParams.sparse_mla_topk_lens;
tllmRunnerParams.ptrSparseMlaTopKLens = mRuntimeSparseAttentionParams.sparse_attn_kv_lens;
tllmRunnerParams.kvPageIdxPtr = reinterpret_cast<KVCacheIndex::UnderlyingType const*>(
mRuntimeSparseAttentionParams.sparse_attn_indices);
if (useDynamicSparseMLA)
Expand Down
4 changes: 2 additions & 2 deletions cpp/tensorrt_llm/kernels/fmhaDispatcher.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -273,11 +273,11 @@ void FmhaDispatcher::run(MHARunnerParams runnerParams)
// The kernel iterates over tokens via numCtasPerSeqQ when maxSeqLenQ > 1.
if (mFixedParams.useTllmGenSparseAttention)
{
bool const useDynamicSparseMLA = runnerParams.sparse_params.sparse_mla_topk_lens != nullptr;
bool const useDynamicSparseMLA = runnerParams.sparse_params.sparse_attn_kv_lens != nullptr;
tllmRunnerParams.mSparseAttention
= useDynamicSparseMLA ? SparseType::DynamicTokenSparse : SparseType::StaticTokenSparse;
tllmRunnerParams.mSparseTopK = runnerParams.sparse_params.num_sparse_topk;
tllmRunnerParams.ptrSparseMlaTopKLens = runnerParams.sparse_params.sparse_mla_topk_lens;
tllmRunnerParams.ptrSparseMlaTopKLens = runnerParams.sparse_params.sparse_attn_kv_lens;
tllmRunnerParams.mKernelType = FmhaKernelType::Generation;
tllmRunnerParams.mUseGenKernelForPrefill = true;
tllmRunnerParams.mMaskType = TrtllmGenAttentionMaskType::Causal;
Expand Down
4 changes: 2 additions & 2 deletions cpp/tensorrt_llm/kernels/sparseAttentionKernels.h
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ struct SparseAttentionParams
// SWA KV pool for dynamic sparse MLA. This is the host KV cache pool pointer for V4 and is
// wired to trtllm-gen's sliding-window KV pool TMA descriptor.
void* sliding_window_kv_cache_pool{nullptr};
int32_t* sparse_mla_topk_lens{nullptr}; // [num_tokens]
int32_t* sparse_attn_kv_lens{nullptr}; // [num_tokens]

int32_t sparse_attn_indices_block_size{1};
int32_t sparse_attn_indices_stride{0};
Expand All @@ -55,7 +55,7 @@ struct SparseAttentionParams
<< "num_sparse_topk: " << this->num_sparse_topk << std::endl
<< "sparse_kv_cache_pool: " << this->sparse_kv_cache_pool << std::endl
<< "sliding_window_kv_cache_pool: " << this->sliding_window_kv_cache_pool << std::endl
<< "sparse_mla_topk_lens: " << this->sparse_mla_topk_lens << std::endl
<< "sparse_attn_kv_lens: " << this->sparse_attn_kv_lens << std::endl
<< "sparse_attn_indices_block_size: " << this->sparse_attn_indices_block_size << std::endl
<< "sparse_attn_indices_stride: " << this->sparse_attn_indices_stride << std::endl;
return ss.str();
Expand Down
4 changes: 2 additions & 2 deletions cpp/tensorrt_llm/nanobind/thop/bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -165,7 +165,7 @@ void initBindings(nb::module_& m)
nb::arg("spec_bl_tree_first_sparse_mask_offset_kv").none(), nb::arg("sparse_kv_indices").none(),
nb::arg("sparse_kv_offsets").none(), nb::arg("sparse_attn_indices").none(),
nb::arg("sparse_attn_offsets").none(), nb::arg("sparse_attn_indices_block_size"),
nb::arg("num_sparse_topk") = std::nullopt, nb::arg("sparse_mla_topk_lens") = std::nullopt,
nb::arg("num_sparse_topk") = std::nullopt, nb::arg("sparse_attn_kv_lens") = std::nullopt,
nb::arg("skip_softmax_threshold_scale_factor_prefill") = std::nullopt,
nb::arg("skip_softmax_threshold_scale_factor_decode") = std::nullopt,
nb::arg("skip_softmax_stat") = std::nullopt, nb::arg("cu_q_seqlens") = std::nullopt,
Expand All @@ -175,7 +175,7 @@ void initBindings(nb::module_& m)
nb::arg("flash_mla_num_splits") = std::nullopt, nb::arg("sage_attn_num_elts_per_blk_q") = 0,
nb::arg("sage_attn_num_elts_per_blk_k") = 0, nb::arg("sage_attn_num_elts_per_blk_v") = 0,
nb::arg("sage_attn_qk_int8") = false, nb::arg("num_contexts") = 0, nb::arg("num_ctx_tokens") = 0,
nb::arg("trtllm_gen_jit_warmup") = false, nb::arg("compressed_kv_cache_pool_ptr") = std::nullopt,
nb::arg("trtllm_gen_jit_warmup") = false, nb::arg("aux_kv_cache_pool_ptr") = std::nullopt,
nb::arg("is_cross") = false, nb::arg("cross_kv") = std::nullopt,
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,
Expand Down
35 changes: 17 additions & 18 deletions cpp/tensorrt_llm/thop/attentionOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -373,13 +373,13 @@ class RunnerBase
torch::optional<torch::Tensor> attention_sinks, torch::optional<torch::Tensor> sparse_kv_indices,
torch::optional<torch::Tensor> sparse_kv_offsets, torch::optional<torch::Tensor> sparse_attn_indices,
torch::optional<torch::Tensor> sparse_attn_offsets, int64_t const sparse_attn_indices_block_size,
int32_t const num_sparse_topk, std::optional<torch::Tensor> sparse_mla_topk_lens,
int32_t const num_sparse_topk, std::optional<torch::Tensor> sparse_attn_kv_lens,
std::optional<torch::Tensor> cu_q_seqlens, std::optional<torch::Tensor> cu_kv_seqlens,
std::optional<torch::Tensor> fmha_scheduler_counter, std::optional<torch::Tensor> mla_bmm1_scale,
std::optional<torch::Tensor> mla_bmm2_scale, std::optional<torch::Tensor> quant_q_buffer,
std::optional<torch::Tensor> flash_mla_tile_scheduler_metadata,
std::optional<torch::Tensor> flash_mla_num_splits, bool trtllm_gen_jit_warmup,
std::optional<int64_t> compressed_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<int64_t> aux_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<torch::Tensor> relative_attention_bias,
std::optional<torch::Tensor> quant_scale_qkv = std::nullopt,
std::optional<torch::Tensor> dsv4_inv_rope_cos_sin_cache = std::nullopt,
Expand Down Expand Up @@ -445,13 +445,13 @@ class Runner : public RunnerBase
torch::optional<torch::Tensor> attention_sinks, torch::optional<torch::Tensor> sparse_kv_indices,
torch::optional<torch::Tensor> sparse_kv_offsets, torch::optional<torch::Tensor> sparse_attn_indices,
torch::optional<torch::Tensor> sparse_attn_offsets, int64_t const sparse_attn_indices_block_size,
int32_t const num_sparse_topk, std::optional<torch::Tensor> sparse_mla_topk_lens,
int32_t const num_sparse_topk, std::optional<torch::Tensor> sparse_attn_kv_lens,
std::optional<torch::Tensor> cu_q_seqlens, std::optional<torch::Tensor> cu_kv_seqlens,
std::optional<torch::Tensor> fmha_scheduler_counter, std::optional<torch::Tensor> mla_bmm1_scale,
std::optional<torch::Tensor> mla_bmm2_scale, std::optional<torch::Tensor> quant_q_buffer,
std::optional<torch::Tensor> flash_mla_tile_scheduler_metadata,
std::optional<torch::Tensor> flash_mla_num_splits, bool trtllm_gen_jit_warmup,
std::optional<int64_t> compressed_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<int64_t> aux_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<torch::Tensor> relative_attention_bias, std::optional<torch::Tensor> quant_scale_qkv,
std::optional<torch::Tensor> dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion) const override
{
Expand Down Expand Up @@ -765,8 +765,8 @@ class Runner : public RunnerBase
op.mRuntimeSparseAttentionParams.sparse_attn_indices_stride
= sparse_attn_indices.has_value() ? sparse_attn_indices.value().size(-1) : 0;
op.mRuntimeSparseAttentionParams.num_sparse_topk = num_sparse_topk;
op.mRuntimeSparseAttentionParams.sparse_mla_topk_lens
= sparse_mla_topk_lens.has_value() ? sparse_mla_topk_lens.value().data_ptr<int32_t>() : nullptr;
op.mRuntimeSparseAttentionParams.sparse_attn_kv_lens
= sparse_attn_kv_lens.has_value() ? sparse_attn_kv_lens.value().data_ptr<int32_t>() : nullptr;
op.mRuntimeSparseAttentionParams.sparse_kv_cache_pool = nullptr;
op.mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool = nullptr;
if (op.mUseSparseAttention && use_kv_cache)
Expand All @@ -775,14 +775,14 @@ class Runner : public RunnerBase
{
auto* kvCachePool = reinterpret_cast<char*>(
host_kv_cache_pool_pointers.value().index({pool_index, 0}).item<int64_t>());
if (sparse_mla_topk_lens.has_value())
if (sparse_attn_kv_lens.has_value())
{
// Deepseek V4 dynamic sparse MLA always uses the SWA pool for now.
op.mRuntimeSparseAttentionParams.sliding_window_kv_cache_pool = kvCachePool;
if (compressed_kv_cache_pool_ptr.has_value())
if (aux_kv_cache_pool_ptr.has_value())
{
op.mRuntimeSparseAttentionParams.sparse_kv_cache_pool
= reinterpret_cast<char*>(compressed_kv_cache_pool_ptr.value());
= reinterpret_cast<char*>(aux_kv_cache_pool_ptr.value());
}
}
else
Expand Down Expand Up @@ -1101,16 +1101,15 @@ void attention(torch::Tensor q, std::optional<torch::Tensor> k, std::optional<to
std::optional<torch::Tensor> sparse_kv_indices, std::optional<torch::Tensor> sparse_kv_offsets,
std::optional<torch::Tensor> sparse_attn_indices, std::optional<torch::Tensor> sparse_attn_offsets,
int64_t const sparse_attn_indices_block_size, std::optional<int64_t> num_sparse_topk,
std::optional<torch::Tensor> sparse_mla_topk_lens,
std::optional<double> skip_softmax_threshold_scale_factor_prefill,
std::optional<torch::Tensor> sparse_attn_kv_lens, std::optional<double> skip_softmax_threshold_scale_factor_prefill,
std::optional<double> skip_softmax_threshold_scale_factor_decode, std::optional<torch::Tensor> skip_softmax_stat,
std::optional<torch::Tensor> cu_q_seqlens, std::optional<torch::Tensor> cu_kv_seqlens,
std::optional<torch::Tensor> fmha_scheduler_counter, std::optional<torch::Tensor> mla_bmm1_scale,
std::optional<torch::Tensor> mla_bmm2_scale, std::optional<torch::Tensor> quant_q_buffer,
std::optional<torch::Tensor> flash_mla_tile_scheduler_metadata, std::optional<torch::Tensor> flash_mla_num_splits,
int64_t sage_attn_num_elts_per_blk_q, int64_t sage_attn_num_elts_per_blk_k, int64_t sage_attn_num_elts_per_blk_v,
bool sage_attn_qk_int8, int64_t num_contexts, int64_t num_ctx_tokens, bool trtllm_gen_jit_warmup,
std::optional<int64_t> compressed_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<int64_t> aux_kv_cache_pool_ptr, bool const is_cross, std::optional<torch::Tensor> cross_kv,
std::optional<torch::Tensor> relative_attention_bias, int64_t relative_attention_max_distance,
std::optional<int64_t> spec_decoding_target_max_draft_tokens, std::optional<torch::Tensor> quant_scale_qkv,
std::optional<torch::Tensor> dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion,
Expand Down Expand Up @@ -1397,10 +1396,10 @@ void attention(torch::Tensor q, std::optional<torch::Tensor> k, std::optional<to
spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset,
spec_decoding_bl_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv, attention_sinks, sparse_kv_indices,
sparse_kv_offsets, sparse_attn_indices, sparse_attn_offsets, sparse_attn_indices_block_size,
num_sparse_topk_value, sparse_mla_topk_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter,
num_sparse_topk_value, sparse_attn_kv_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter,
mla_bmm1_scale, mla_bmm2_scale, quant_q_buffer, flash_mla_tile_scheduler_metadata, flash_mla_num_splits,
trtllm_gen_jit_warmup, compressed_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias,
quant_scale_qkv, dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion);
trtllm_gen_jit_warmup, aux_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias, quant_scale_qkv,
dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion);
}

if ((num_generations > 0) && (attn_input_type != AttentionInputType::ContextOnly))
Expand All @@ -1420,10 +1419,10 @@ void attention(torch::Tensor q, std::optional<torch::Tensor> k, std::optional<to
spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, spec_decoding_bl_tree_mask_offset,
spec_decoding_bl_tree_mask, spec_bl_tree_first_sparse_mask_offset_kv, attention_sinks, sparse_kv_indices,
sparse_kv_offsets, sparse_attn_indices, sparse_attn_offsets, sparse_attn_indices_block_size,
num_sparse_topk_value, sparse_mla_topk_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter,
num_sparse_topk_value, sparse_attn_kv_lens, cu_q_seqlens, cu_kv_seqlens, fmha_scheduler_counter,
mla_bmm1_scale, mla_bmm2_scale, quant_q_buffer, flash_mla_tile_scheduler_metadata, flash_mla_num_splits,
trtllm_gen_jit_warmup, compressed_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias,
quant_scale_qkv, dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion);
trtllm_gen_jit_warmup, aux_kv_cache_pool_ptr, is_cross, cross_kv, relative_attention_bias, quant_scale_qkv,
dsv4_inv_rope_cos_sin_cache, enable_dsv4_epilogue_fusion);
}

TLLM_LOG_TRACE("Attention op stops at layer %d", local_layer_idx);
Expand Down
5 changes: 2 additions & 3 deletions cpp/tensorrt_llm/thop/attentionOp.h
Original file line number Diff line number Diff line change
Expand Up @@ -80,8 +80,7 @@ void attention(torch::Tensor q, std::optional<torch::Tensor> k, std::optional<to
std::optional<torch::Tensor> sparse_kv_indices, std::optional<torch::Tensor> sparse_kv_offsets,
std::optional<torch::Tensor> sparse_attn_indices, std::optional<torch::Tensor> sparse_attn_offsets,
int64_t const sparse_attn_indices_block_size, std::optional<int64_t> num_sparse_topk,
std::optional<torch::Tensor> sparse_mla_topk_lens,
std::optional<double> skip_softmax_threshold_scale_factor_prefill,
std::optional<torch::Tensor> sparse_attn_kv_lens, std::optional<double> skip_softmax_threshold_scale_factor_prefill,
std::optional<double> skip_softmax_threshold_scale_factor_decode, std::optional<torch::Tensor> skip_softmax_stat,
std::optional<torch::Tensor> cu_q_seqlens, std::optional<torch::Tensor> cu_kv_seqlens,
std::optional<torch::Tensor> fmha_scheduler_counter, std::optional<torch::Tensor> mla_bmm1_scale,
Expand All @@ -90,7 +89,7 @@ void attention(torch::Tensor q, std::optional<torch::Tensor> k, std::optional<to
std::optional<torch::Tensor> flash_mla_num_splits = std::nullopt, int64_t sage_attn_num_elts_per_blk_q = 0,
int64_t sage_attn_num_elts_per_blk_k = 0, int64_t sage_attn_num_elts_per_blk_v = 0, bool sage_attn_qk_int8 = false,
int64_t num_contexts = 0, int64_t num_ctx_tokens = 0, bool trtllm_gen_jit_warmup = false,
std::optional<int64_t> compressed_kv_cache_pool_ptr = std::nullopt, bool const is_cross = false,
std::optional<int64_t> aux_kv_cache_pool_ptr = std::nullopt, bool const is_cross = false,
std::optional<torch::Tensor> cross_kv = std::nullopt,
std::optional<torch::Tensor> relative_attention_bias = std::nullopt, int64_t relative_attention_max_distance = 0,
std::optional<int64_t> spec_decoding_target_max_draft_tokens = std::nullopt,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -220,7 +220,7 @@ Figure 4 illustrates the prediction implementation within TensorRT LLM. To suppo

**Auxiliary memory management.** Managing the paged KT cache presented another challenge. `RocketKVCacheManager` inherits from `KVCacheManager` and extends it with a dedicated `BlockManager` for the auxiliary KT cache at the Python level. The main KV cache and the KT cache share block IDs for each request, so that the lifecycle of KT cache blocks is automatically tied to the corresponding KV cache blocks. The `BlockManager` handles slot allocation and deallocation for the KT cache independently, while `RocketKVCacheManager` overrides methods such as `get_cache_bytes_per_token` and `prepare_resources` to ensure that memory sizing accounts for the extra KT cache footprint and that the correct KT cache pointers are passed to prediction kernels at each step. This design keeps the integration lightweight and easy to iterate on, though it inherits the limitations of Python-level management—namely, no automatic support for KV cache reuse or disaggregated serving.

The concrete implementation can be found in `tensorrt_llm/_torch/attention_backend/sparse/rocket.py`.
The concrete implementation can be found in `tensorrt_llm/_torch/attention_backend/sparse/rocket/`.

### DeepSeek Sparse Attention (DSA)

Expand Down Expand Up @@ -261,7 +261,7 @@ As with RocketKV, a dedicated metadata class `DSATrtllmAttentionMetadata` is def

**Auxiliary memory management.** DSA requires an auxiliary **indexer K cache** to store the low-rank K projections for reuse across decoding steps. `DSAKVCacheManager` inherits from `KVCacheManager`, but unlike RocketKV's Python-level KT cache management, DSA's indexer K cache is integrated directly into the C++ `KVCacheManager`. This design enables compatibility with advanced features such as KV cache reuse, chunked prefill, and disaggregated serving—features that would be difficult to support with a Python-level manager.

The concrete implementation can be found in `tensorrt_llm/_torch/attention_backend/sparse/dsa.py`.
The concrete implementation can be found in `tensorrt_llm/_torch/attention_backend/sparse/dsa/`.

For a comprehensive description of DSA kernel optimizations, precision strategies, feature support (MTP, disaggregated serving, Wide-EP), and benchmark results, please refer to the dedicated blog post: [Optimizing DeepSeek-V3.2 on NVIDIA Blackwell GPUs](blog15_Optimizing_DeepSeek_V32_on_NVIDIA_Blackwell_GPUs.md).

Expand Down
Loading
Loading