DeepSeek Sparse Attention (DSA)#
This is an experimental API and subject to change.
Overview#
The DeepSeek Sparse Attention (DSA) module integrates a set of CuTe-DSL
kernels that support the sparse-attention path used by DeepSeek-style models.
Most kernels target Hopper (SM90) and Blackwell (SM100+) GPUs. Dense Indexer
Forward and standalone Indexer Top-K support both architectures; the combined
compressed-logits + Top-K path is SM100-only. The kernels are delivered as
Python classes / wrappers that follow the same APIBase pattern as other
cuDNN Frontend operations.
Scope: this module ships CuTe-DSL kernels for DSA backward, indexer
scores/top-K, sparse/dense score recompute, and sparse/dense indexer
backward. The production
sparse-attention forward kernel (FlashMLA) is C++ and is not integrated
here; when evaluating the backward, use the pure-PyTorch reference in
test/python/fe_api/dsa/dsa_reference.py::ref_sparse_attention_forward.
The module packages the following operations:
Sparse Attention Backward – DSA backward (FlashMLA-shape, SM90/SM100).
Indexer Forward – CuTe-DSL score kernel (Q @ K^T, ReLU, head reduce, ratio causal mask) that materializes dense scores.
Combined Indexer Forward + Top-K – SM100 compact score generation, Top-K selection, and optional Top-K softmax in one public API call.
Indexer Top-K – SM90+ CuTe-DSL radix top-K kernel with per-row
seq_lens.Sparse Indexer / Attention Score Recompute – sparse (top-K) recompute of indexer and attention scores for training loss.
Dense Indexer / Attention Score Recompute – dense (full-KV) analogues of the above.
Indexer Backward – three-stage pipeline (score-grad, three GEMMs, dtype cast) for sparse top-K score tensors.
Dense Indexer Backward – full-KV counterpart of Indexer Backward.
Architecture#
Q, K, W ──┬─► IndexerForward ──► scores ──► IndexerTopK ──► topk_idxs
└─► IndexerForwardTopK ──► topk_idxs, logits, predict
│
v
[FlashMLA fwd — external, C++] ──► out, lse
│
dout ──┤
v
SparseAttentionBackward
│
v
dq, dkv, d_sink
Training-score loss path:
attn_score, index_score ──► IndexerBackward ──► d_index_q, d_weights, d_index_k
(SparseIndexer/AttnScoreRecompute and DenseIndexer/AttnScoreRecompute
produce these score tensors; DenseIndexerBackward consumes dense raw scores.)
Installation#
pip install nvidia-cudnn-frontend[cutedsl]
API Usage#
DSA Namespace#
from cudnn import DSA
DSA.SparseAttentionBackward
DSA.sparse_attention_backward_wrapper
DSA.IndexerForward
DSA.indexer_forward_wrapper
DSA.indexer_forward_top_k_wrapper
DSA.compress_topk_cand_buffer_size
DSA.compress_topk_cand_buffer_size_thd
DSA.IndexerTopK
DSA.indexer_top_k_wrapper
DSA.SparseIndexerScoreRecompute
DSA.sparse_indexer_score_recompute_wrapper
DSA.SparseAttnScoreRecompute
DSA.sparse_attn_score_recompute_wrapper
DSA.DenseIndexerScoreRecompute
DSA.dense_indexer_score_recompute_wrapper
DSA.DenseAttnScoreRecompute
DSA.dense_attn_score_recompute_wrapper
DSA.IndexerBackward
DSA.indexer_backward_wrapper
DSA.DenseIndexerBackward
DSA.dense_indexer_backward_wrapper
Components#
1. Sparse Attention Backward#
Backward pass for DeepSeek Sparse Attention. Expects the forward outputs
(out, lse) from FlashMLA (or the PyTorch reference).
Inputs
q:(total_S_q, H, D)BF16/FP16kv:(total_S_kv, D)(K = V; MQA)out,dout:(total_S_q, H, D_v)lse:(total_S_q, H)FP32attn_sink:(H,)FP32topk_idxs:(total_S_q, topk_max)INT32 (global)topk_length(optional):(total_S_q,)INT32 — per-query valid count
On SM100, the public backward entry point automatically selects the tuned
kernel from q.shape[1:3]: H16 with head_dim=576 uses the dedicated M128
sparse-row pipeline, while head_dim=512, H32/H64, and other supported shapes
use the generic M64 pipeline. No backend or tile-size argument is required.
SM90 continues to use its Hopper-specific implementation.
Outputs — tuple
(dq, dkv, d_sink)Constraints — SM90 or SM100; SM90 supports the FlashMLA DSA shape with
head_dim ∈ {512, 576}
result = DSA.sparse_attention_backward_wrapper(
q, kv, out, dout, lse, attn_sink, topk_idxs,
softmax_scale=1.0 / math.sqrt(D),
topk_length=topk_length,
)
dq, dkv, d_sink = result["dq"], result["dkv"], result["d_sink"]
2. Indexer Forward#
Computes dense indexer scores:
S[b, q, k] = sum_h ReLU(Q_h · K_h^T) · W_h with a ratio-causal mask.
For local query row q_local, valid KV columns satisfy
k_local < clamp((q_causal_offsets[b] + q_local + 1) // ratio, 0, seqlen_k_b).
When q_causal_offsets is omitted, all offsets are zero.
q_causal_offsets[b] is the global uncompressed token index corresponding to
local q[0] for batch or THD segment b. It is not the packed storage offset
from cu_seqlens_q: cu_seqlens_q locates where a local Q segment is stored,
while q_causal_offsets locates that segment in the global causal timeline.
The K columns are assumed to be a compressed-KV prefix starting at global
compressed column 0.
Inputs
q:(B, S_q, H_q, D)BF16, or the architecture-specific FP8 format described below.k:(B, S_k, H_kv, D)BF16, or the architecture-specific FP8 format described below.w:(B, S_q, H_q)BF16. The SM90 FP8 path also accepts FP32 when weights have already been pre-scaled byq_scale * sm_scale.q_causal_offsets(optional): CUDA INT32 tensor with one entry per batch/THD segment, on the same device asq.
Output —
scores:(B, S_q, S_k)FP32.Precision paths
SM90
precision="fp8": Q/K use E4M3 andq_scale/k_scaleare FP32 descales with one value per token/head. Setreturn_lse=True(or providelse_out) to compute LSE in the same kernel invocation.SM100
precision="mxfp8": Q/K use E4M3 with block-scaled, packed E8M0 scale tensors;sf_vec_sizeis currently fixed at 32 andqhead_per_kv_head ∈ {32, 64}. THD inputs additionally requirecu_seqlens_q_scale_padded/cu_seqlens_k_scale_padded: contiguous CUDA INT32 prefix tensors of shape(B + 1,)on the Q/K device. The interface deliberately performs no device-to-host copy to inspect their values. The caller guarantees that each device-side prefix starts at zero, is monotonic, covers the corresponding logical sequence, satisfies the packed-row alignment (each Q span timesqhead_per_kv_headis a multiple of 128 rows; each K span is a multiple of 128 tokens), and stays within the packed scale storage.
Constraints —
head_dim == 128. The SM90 direct path supportsqhead_per_kv_head ∈ {16, 32, 64}and currently requiresH_kv == 1; SM100 BF16 dense and combined Top-K paths supportqhead_per_kv_head ∈ {32, 64}, as do their MXFP8 paths. All currently requireH_kv == 1(MQA).
result = DSA.indexer_forward_wrapper(
q, k, w, ratio=4, q_causal_offsets=q_causal_offsets,
)
scores = result["scores"]
3. Combined Indexer Forward + Top-K#
DSA.indexer_forward_top_k_wrapper provides a one-call alternative to the
two-call dense indexer_forward_wrapper + indexer_top_k_wrapper sequence
when compact score generation is desired. It returns aligned indices INT32
and logits FP32, plus softmax FP32 by default, without materializing dense
scores. Pass
return_softmax=False to omit softmax. BSHD output shape is
(B, S_q, top_k); THD output shape is (total_q, top_k). Padded slots are
-1/-inf. The selected set is Top-K, but its slot order is not guaranteed to
be descending. Set deterministic=True to break exact-value ties at the K-th
boundary toward the smallest local KV indices, making the selected set
reproducible across launches. This does not sort the output slots; the default
False path retains the faster scheduling-dependent tie-break.
The combined compressed path is SM100-only. Both BSHD and THD support
BF16 and MXFP8. topk_indices_global=True is the default. Optional caller-owned
candidate/output/softmax/LSE buffers avoid per-call allocations; size the
candidate buffer with compress_topk_cand_buffer_size for BSHD or
compress_topk_cand_buffer_size_thd for THD. LSE is supported for BSHD and THD
with both BF16 and MXFP8. Explicit microbatching is BF16 BSHD-only and cannot
be combined with LSE or explicit q_causal_offsets. BSHD with explicit
q_causal_offsets currently computes its per-batch candidate offsets eagerly
and is not CUDA-graph-capturable; THD capture requires the caller-provided
offsets and buffers returned by compress_topk_cand_buffer_size_thd. Compact
addressing requires 0 <= q_causal_offsets[b]; rows extending beyond the KV
prefix are clamped to seqlen_k_b. THD MXFP8 uses the same caller-guaranteed
device-side scale-prefix contract described in §2; the interface does not
validate prefix values with a device-to-host copy.
cand_floats = DSA.compress_topk_cand_buffer_size(
B, S_q, S_k, ratio=4, return_lse=True,
)
cand = torch.empty(cand_floats, dtype=torch.float32, device=q.device)
result = DSA.indexer_forward_top_k_wrapper(
q, k, w, top_k=512,
ratio=4,
cand_buffer=cand,
return_lse=True,
)
topk_indices, topk_logits = result["indices"], result["logits"]
predict = result["softmax"]
lse = result["lse"]
4. Indexer Top-K#
Radix top-K kernel for selecting candidate KV indices from indexer scores, with variable per-row effective length.
Inputs
input_values:(n_rows, num_cols)FP32/FP16/BF16seq_lens:(batch_size,)INT32 (per-batch effective column count)
Outputs — tuple
(indices, values)(values isNonewhenreturn_val=False)Constraints — SM90+,
top_k ≤ 2048
result = DSA.indexer_top_k_wrapper(
scores.reshape(-1, scores.shape[-1]),
seq_lens, top_k=512,
)
indices, values = result["indices"], result["values"]
5. Sparse Indexer Score Recompute#
Computes softmax over top-K entries of the indexer score:
predict[b, q, i] = softmax_i(sum_h ReLU(Q_h · K_{topk[i]}^T) · W_h).
Inputs:
q_indexer,k_indexer,weights,topk_indices(optionaltopk_length).topk_indicesare per-batch local KV ids by default; passtopk_indices_global=Truewhen using ids encoded asbatch_idx * S_k + local_idx.Output —
predict:(B, S_q, topk)FP32.
6. Sparse Attn Score Recompute#
L1-normalised head-summed softmax over top-K entries:
target[b, q, i] = sum_h exp(Q_h · K_{topk[i]}^T · scale - LSE_h) / Z.
Inputs:
q_attn,k_attn,lse,topk_indices,softmax_scale(optionaltopk_length).topk_indicesare per-batch local KV ids by default; passtopk_indices_global=Truewhen using ids encoded asbatch_idx * S_k + local_idx.Output —
target:(B, S_q, topk)FP32.Note: the wrapper handles the
-log2(e) * lsepreprocessing internally.
7. Dense Indexer / Dense Attn Score Recompute#
Full-KV (no top-K) analogues of §5 and §6. Each returns {'out', 'denom'}.
They apply the same ratio-causal mask as Indexer Forward; masked positions are
written as -inf and excluded from denom. Pass the same q_causal_offsets to
all dense score tensors that feed the same loss path.
On SM100, Indexer Forward and Dense Indexer Score Recompute use the same
unified kernel implementation: forward runs it with compute_lse=False, while
dense indexer score recompute runs it with compute_lse=True. The shared
implementation lives in score_recompute; indexer_forward only imports it.
Dense Attention Score Recompute has a separate MXFP8 kernel because its score
and normalization semantics differ from the indexer path.
8. Indexer Backward#
Three-stage sparse top-K pipeline that produces the training gradients for the indexer tower:
ScoreGradSm90/ScoreGradSm100(kernel 1) — in-place score-grad precompute fromattn_score(target) andindex_score(predict).IndexerBackwardSm90/IndexerBackwardSm100(kernel 2) — three warp-specialised GEMMs produced_index_q,d_weights, and adIndexK_f32accumulator.Pure-torch dtype cast (kernel 3) converts
dIndexK_f32to the output dtype.
The TileLang fallback present in the upstream repo is dropped here
(CuTe-DSL only). If the CuTe-DSL path fails the wrapper raises
RuntimeError rather than silently falling back.
grad_loss = torch.ones((), dtype=torch.float32, device=index_q.device)
result = DSA.indexer_backward_wrapper(
index_q, weights, index_k,
attn_score, index_score, topk_indices,
grad_loss=grad_loss, sm_scale=1.0, loss_coeff=1.0, block_I=128,
)
d_index_q, d_weights, d_index_k = (
result["d_index_q"], result["d_weights"], result["d_index_k"],
)
When compressed forward returns its fused softmax, backward can skip both
indexer Q@K score recompute and the separate logits softmax. Pass softmax
directly as index_score; backward consumes and overwrites this buffer, so
pass softmax.clone() if it must be preserved. attn_score must use the same
valid-slot mask. Because compressed forward returns global indices by default,
also pass topk_indices_global=True unless forward used
topk_indices_global=False. The public sparse indexer_backward_wrapper has a
BSHD-shaped interface; BF16 THD tensors can use zero-copy B=1 views (squeeze
the singleton K head and add a batch dimension) together with global Top-K
indices. FP8 and MXFP8 indexer backward are not currently supported because
the backward wrapper requires BF16 Q/K/W inputs.
SM100 sparse backward v2 (opt-in)#
backend="sm100_v2" on IndexerBackward / indexer_backward_wrapper
selects an alternative kernel-2 (GEMM stage) implementation. backend is a
string enum, one of {"default", "sm100_v2"} (default "default"; an
unknown value raises ValueError). The semantics are
request-or-fail: outside the envelope below the wrapper raises
ValueError (or RuntimeError off SM100) and never silently falls back to
the default backend.
What it computes differently — weights are upcast to fp32 in-register (exact) and the fp32 per-slot gradient matrix is split into a two-term BF16 expansion (
hi + lo) before the MMAs, so each individual product in the dQ/dK contractions is exact in the fp32 accumulator when it is finite there (~675x lower gradient-matrix representation error than the default single-BF16 rounding, measured as an aggregate over 1e6 random values; individual values gain less when theirloterm is small).lois ~2^-9 ofhi, so it starts to underflow for gradient-matrix magnitudes below ~1e-35 (measured 518x at 1e-35, 63x at 1e-36, and no advantage left at 1e-38), and the FP32->BF16 conversion is non-saturating, so a magnitude at or above BF16’s round-to-nearest overflow threshold (2^128 - 2^119~= 3.3962e38, itself above BF16’s 3.3895e38 maximum) becomes an infinity inhiand the opposite infinity inlo, which sum to NaN rather than being clamped.d_weightsis reduced deterministically —d_weightsandd_index_qare bitwise run-to-run stable;d_index_kstays in the same fp32-atomic summation-order class as the default backend.Output dtype selects output precision —
d_weightsandd_index_kaccept FP32 buffers, which receive the fp32 accumulators directly; the headline accuracy gains require caller-supplied FP32 buffers. The default wrapper-allocated outputs keep the input dtypes (BF16), which rounds the extra accuracy back to the BF16 representation floor (~1.7e-3 relative). FP32d_index_kis zeroed internally (no caller pre-zero contract).d_index_qis always BF16.Envelope (validated by
check_supportwhen a plan is created, and re-validated against the real tensors on every execute – cache hits included – before any buffer is mutated) — SM100 (10, 0) only;H == 64,D == 128,block_I == 128;topk % 128 == 0with128 <= topk <= 2048;sm_scale > 0; contiguous same-device tensors;attn_score/index_score/topk_indicesshaped(B, S_q, topk).sm_scalefolds into the gradient in kernel 2 while the relu gate reads the unscaled score, which is equivalent for any positive scale except at the underflow-to-zero boundary: for a slot whose scaled scoresm_scale * Srounds to zero the default backend drops the slot and this backend keeps it, and because the gate is a step function that slot’s full dQ/dK contribution (g' * w, not scaled byS) is what differs. The default backend’s multiply preserves subnormals, so reaching it takes an exactsm_scale * Sbelow2^-150(~7.0e-46).Scratch / workspace behavior —
attn_score/index_scoreare consumed in place exactly like the default backend (attn_scoreis left holding kernel 1’sgrad_signal;sm_scalefolds inside kernel 2 without touching the buffer). The backend owns one piece of per-plan workspace — the dynamic-ticket counter — allocated on first execute and reused by every later one. A BF16d_index_kadditionally needs aB * S_k * Dfp32 accumulator (2 MiB = 2,097,152 bytes at B=1, S_k=4096, D=128, growing withS_k); that one comes from PyTorch’s caching allocator on every call, which is a pool hit in steady state — it has to be re-zeroed per call anyway, and this way it stays reclaimable viatorch.cuda.empty_cache()instead of staying pinned for as long as the plan is cached. The wrapper still allocates any output buffer you do not pass in.Concurrency — executions sharing one plan must not overlap on the device (the ticket counter is per-plan workspace). One plan serves one device: the workspace is device-resident, and execution rejects tensors on any other device before touching the score buffers. The wrapper keys its plan cache on the CUDA device and on the resolved stream (
streamwhen given, otherwisetorch.cuda.current_stream()at call time), so calls that differ in device or stream get a private plan and private workspace; calls that land on the same cache entry must not be allowed to overlap. The key is the integer stream handle, plus the calling thread’s id for the one handle CUDA does not make unique across host threads:cudaStreamPerThreadis the value 2 in every thread and means “the calling thread’s own stream”, so two threads that pass it explicitly get a private plan each. Users drivingIndexerBackwardobjects directly must use one object per device, and one per stream wherever those streams’ executions can overlap; a single object may target any stream if its executions are serialized.Local top-k ids (
topk_indices_global=False) are masked against the per-batchS_kbefore the batch offset is applied, in-kernel: ids< 0or>= S_kcontribute nothing and can never alias a neighbouring batch.
# fp32 output buffers unlock the full accuracy gain
d_weights = torch.empty(B, S_q, H, dtype=torch.float32, device="cuda")
d_index_k = torch.empty(B, S_k, D, dtype=torch.float32, device="cuda")
result = DSA.indexer_backward_wrapper(
index_q, weights, index_k,
attn_score, index_score, topk_indices,
grad_loss=grad_loss, sm_scale=1.0, loss_coeff=1.0, block_I=128,
backend="sm100_v2",
d_weights=d_weights, d_index_k=d_index_k,
)
9. Dense Indexer Backward#
Full-KV counterpart to Indexer Backward. It consumes raw dense score tensors and denominators produced by Dense Indexer / Dense Attn Score Recompute.
Inputs
index_q:(B, S_q, H, D)BF16weights:(B, S_q, H)BF16index_k:(B, S_k, D)BF16attn_score,index_score:(B, S_q, S_k)FP32 raw dense scoresattn_l1norm,index_lse:(B, S_q)FP32 denominatorsq_causal_offsets(optional): same offsets used for the corresponding Dense Indexer / Dense Attn Score Recompute outputs.
Outputs —
d_index_q,d_weights,d_index_kConstraints — SM90 or SM100+,
H >= 64,ratio >= 1
grad_loss = torch.ones((), dtype=torch.float32, device=index_q.device)
dense_index = DSA.dense_indexer_score_recompute_wrapper(
index_q, index_k.unsqueeze(2), weights,
q_causal_offsets=q_causal_offsets,
)
dense_attn = DSA.dense_attn_score_recompute_wrapper(
attn_q, attn_k, lse, softmax_scale,
q_causal_offsets=q_causal_offsets,
)
result = DSA.dense_indexer_backward_wrapper(
index_q, weights, index_k,
dense_attn["out"], dense_attn["denom"],
dense_index["out"], dense_index["denom"],
grad_loss=grad_loss, sm_scale=1.0, loss_coeff=1.0, block_I=128, ratio=1,
q_causal_offsets=q_causal_offsets,
)
Limitations#
Architecture support — Sparse Attention Backward, Score Recompute, Indexer Forward, Indexer Top-K, and Indexer Backward support SM90 and SM100. The combined compressed-logits + Top-K forward is SM100-only; the standalone Indexer Top-K remains SM90+.
No fused forward — the production forward is FlashMLA (C++); this module ships only the CuTe-DSL kernels.
Indexer Forward only supports
head_dim = 128. SM90 supportsqhead_per_kv_head ∈ {16, 32, 64}withH_kv = 1; SM100 BF16 and MXFP8 supportqhead_per_kv_head ∈ {32, 64}. Both the dense and combined Top-K paths requireH_kv = 1.Standalone Top-K only up to 2048;
top_k > 2048is not supported by its radix Top-K kernel. The combined compressed path uses a separate stage-2 implementation.Compressed-path limits — the stage-1 compact score kernel is MQA-only (
H_kv = 1); MXFP8 requiresqhead_per_kv_head ∈ {32, 64}; explicit microbatching cannot be combined with MXFP8, LSE, or explicit per-batch causal offsets.