DeepSeek Sparse Attention (DSA)
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. Sparse-attention
forward and the combined compressed-logits + Top-K path are SM100-only;
sparse-attention forward is limited to the mapped SM100-family capabilities
10.0, 10.3, and 10.7. Dense Indexer Forward and standalone Indexer Top-K
support both architectures. 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 sparse-attention forward and backward, indexer scores/top-K, sparse/dense score recompute, and sparse/dense indexer backward. Sparse forward covers regular H64 and small-top-k Prefill H128; it does not include the H128 decode/split-KV mode.
The module packages the following operations:
- Sparse Attention Forward – sparse Prefill MQA for the supported SM100 H64/H128 shapes.
- Sparse Attention Backward – DSA backward for flat MQA tensors on 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
Installation
API Usage
DSA Namespace
Components
1. Sparse Attention Forward
Sparse Prefill forward for flat MQA tensors. For query row t, head h, and
valid top-k slot j, the score is
softmax_scale * dot(q[t,h,:], kv[topk_idxs[t,j],:]). Invalid and OOB indices
have score -inf; duplicate indices remain distinct softmax slots. The
attention sink participates only in the output denominator, not in any stats.
- Inputs
q:(total_S_q, H, D_qk)FP16 or BF16kv:(total_S_kv, D_qk), with the same dtype asq(K = V latent; MQA)topk_idxs:(total_S_q, K)INT32 global token ids; K may be any nonnegative value and is padded internally to a multiple of 64attn_sink(optional):(H,)FP32topk_length(optional):(total_S_q,)INT32, clamped to[0, K]softmax_scale(optional runtime float): defaults to1 / sqrt(D_qk)indexer_topk: compile-time prefix length
- Outputs — stable
TupleDictkeys:out:(total_S_q, H, 512), with the same dtype asqmax_logits:(total_S_q, H)FP32lse:(total_S_q, H)FP32, KV-only and excluding the sinklse_indexer:(total_S_q, H)FP32, orNoneforindexer_topk=0
- Variants
- H64:
D_qk in {512, 576},indexer_topk in {0, 512, 1024, 2048} - H128 small-top-k Prefill:
D_qk=512,indexer_topk in {0, 512, 1024}
- H64:
- Indexer prefix constraint —
indexer_topk <= K; whentopk_lengthis supplied and indexer LSE is enabled, every row must havetopk_length >= indexer_topk, matching the tile-boundary snapshot contract. - Empty rows —
out=0,max_logits=-inf,lse=+inf; an enabledlse_indexeris also+inf. - Compilation lifecycle —
SparseAttentionForward.compile()validates and primes the API object, while the concrete CuTe kernel is JIT-compiled on its firstexecute()/wrapper call because padding and live DLPack layouts are runtime properties. - Layout and stream — non-contiguous or under-aligned Q/KV inputs are
materialized on the selected stream. An explicit
streammust be acuda.bindings.driver.CUstreambelonging toq.device.
For D_qk=576, QK uses all 576 dimensions while V/output uses only the first
512 KV dimensions. attn_sink=None is equivalent to an all--inf sink.
2. Sparse Attention Backward
Backward pass for DeepSeek Sparse Attention. Expects the forward wrapper’s
out and KV-only lse (or equivalent tensors).
- 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. The H128 two-CTA backends (D512 and D576) clamp it to[0, topk_max]and ignore every slot after that prefix, including valid indices pointing to nonfinite KV rows; the other backends expecttopk_length <= topk_max. Omitting this tensor uses alltopk_maxslots.
On Blackwell SM100/SM103, the public backward entry point automatically selects
the tuned kernel from the device, dtype, and tensor shape. On SM100 (10, 0) and
SM103 (10, 3) devices, BF16 H128 with head_dim = head_dim_v = 512 and
topk_max ∈ {128, 512, 1024, 1152, 2048} uses the two-CTA specialization.
Contiguous BF16 H128 with head_dim = 576, head_dim_v = 512, and the same
topk_max set uses the H128/D576 two-CTA specialization. H16 with
head_dim=576 uses the dedicated M128 sparse-row pipeline. FP16, other head
counts and dimensions, every other topk_max, and noncontiguous H128/D576
inputs retain the existing generic/H16/H32 selection. Other compute
capabilities, including SM107, do not select the two-CTA paths. No backend or
tile-size argument is required. SM90 continues to use its Hopper-specific
implementation.
The H128 specialization keeps the five tensor-core products in one two-CTA main kernel. It publishes FP32 O-dot-dO and folded-LSE statistics to the caller-provided scratch workspace, converts the FP32 dKV workspace to the public BF16 output, and completes dSink with a separate FP32 reduction kernel. The helper launches do not change the two-CTA topology of the core computation.
The H128/D576 specialization applies the same two-CTA topology to the MLA
layout (head_dim = 576, head_dim_v = 512). Its compiled sequence
initializes the FP32 dKV workspace and the dSink output on the launch stream,
accumulates dKV in that workspace, and converts it to the public BF16 output
with a helper kernel. Short query batches split each query’s top-k range
across otherwise idle two-SM clusters; the split count is fixed when the plan
is compiled.
On SM100 with H16/H32/H64/H96/H128, deterministic=True selects a bounded-wave
M64 implementation. Queries run in same-stream waves of 128 CTAs; CTA lane
i is the sole writer of FP32 dKV shard i in each wave. H16/H32 use a masked
M64 head tile. H96/H128 split each query wave into ordered M64 head-block
launches, preserving single-writer ownership without multiplying the number
of shards. Kernel launch ordering serializes both head blocks and shard reuse
across waves, so the protocol needs neither semaphores nor cooperative launch.
A fixed-order two-stage reduction combines the 128 shards, and one CTA per
head reduces d_sink. The additional dKV workspace is
128 * round_up(total_S_kv, 8) * round_up(D, 8) * sizeof(float) bytes. Use
scratch_workspace_bytes() as the authoritative full scratch size.
deterministic=True takes precedence over the two-CTA selection: the BF16
H128/D512 and H128/D576 envelopes also run the bounded-wave M64 kernel when
determinism is requested, because the two-CTA paths accumulate dKV with FP32
atomics.
Treat this as a reproducibility requirement rather than a performance-tuning
knob: keep the default False when bitwise run-to-run stability is not needed.
SparseAttentionBackward.scratch_workspace_bytes() reports the full SM100
scratch requirement. Pass a contiguous CUDA uint8 tensor of at least this
size to execute(..., workspace=workspace) and reuse it across calls; the
compiled kernel initializes the dKV accumulator on every execution. The
high-level wrapper accepts the same optional workspace= argument and only
allocates convenience scratch when it is omitted. The H128/D576 two-CTA plan
compiles its kernel in compile(), and its execute() additionally requires
caller-provided dq, dkv, and d_sink buffers: it never allocates or
compiles during execution. The wrapper allocates those outputs when they are
omitted. Other backends do not accept a caller-provided d_sink.
- Outputs — tuple
(dq, dkv, d_sink) - Constraints — SM90 or Blackwell SM100/SM103; SM90 supports flat MQA tensors with
head_dim ∈ {512, 576}
3. Indexer Forward (score-only)
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.
- SM90
- 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).
4. 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 §3; the interface does not
validate prefix values with a device-to-host copy.
Leave microbatch_rows=-1 to use the BSHD long-sequence memory policy; it
enables windowing only for compatible BF16 shapes. Pass 0 to force one
launch, or a positive value only when deliberately controlling the scratch
footprint, and pass the same value to compress_topk_cand_buffer_size.
Global versus local ids is primarily an interoperability and allocation
choice, not an Indexer Backward tuning choice. Keep the default global ids
when a downstream consumer requires flat KV ids. If every consumer accepts
local ids, prefer topk_indices_global=False. Avoiding the local-to-global
conversion can slightly improve wrapper performance, particularly for short
sequences. Local ids also avoid conversion temporaries and enable a
strictly zero-extra-allocation captured path when all documented buffers are
supplied. The global-id conversion is CUDA-graph-safe; during capture, its
temporaries are recorded in the graph’s private pool. Include any downstream
conversion when comparing end-to-end performance.
5. 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). Usereturn_val=Falsewhen only the indices are consumed, so no values output buffer is allocated or written. - Constraints — SM90+,
top_k ≤ 2048
Compact vs. non-compact sparse indices
Providing topk_length declares a compact input layout for Sparse Attention
Backward and the sparse score-recompute kernels: in each row, the first
topk_length[row] slots must contain valid ids and every later slot must be
filled with -1. The padding value is required even when topk_length is
provided, because the SM100 attention score-recompute wrapper may select a
non-compact kernel that identifies invalid slots by their index values. Use
compactify_wrapper first when invalid entries are interspersed. A 3-D input
is returned flattened as (B * S_q, topk) with a (B * S_q,) length tensor;
reshape both when a BSHD score-recompute wrapper expects batched shapes. The
operation preserves the input id convention; it does not convert local ids to
global ids.
There is no universal faster representation. Compact execution can skip
sparse tiles when valid prefixes are substantially shorter than the physical
top-k width, while non-compact execution can benefit from static loop bounds
when rows are nearly full. The SM100 sparse attention score-recompute wrapper
uses separately tuned tile policies and can internally select its non-compact
path when that has better code generation. If conversion is required solely
for performance, benchmark the entire compactify_wrapper + consumer sequence
for the target length distribution. Indexer Backward has no topk_length
argument and always processes its fixed slot width; its score tensors and ids
must remain aligned slot-for-slot.
For BF16 workloads with causal sparse indices, the following are useful starting points; the best choice still depends on the shape and valid-length distribution:
- On SM100, prefer trying the non-compact path for Sparse Attention Score Recompute when optimizing this consumer in isolation; it can be slightly faster. For Sparse Indexer Score Recompute, preserve the producer’s representation, since the performance difference is generally small.
- On SM90, prefer compact inputs for both sparse score-recompute kernels when the valid-prefix length is already available; compact execution can be slightly faster.
- For Sparse Attention Backward on either architecture, prefer compact inputs when the valid-prefix length is already available. The benefit may be small when most rows are full.
Retain local ids when the producer and every downstream consumer accept them.
If compactify_wrapper or an id conversion must be added, time that operation
together with the consumer before choosing a representation.
6. 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.
7. 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.
8. Dense Indexer / Dense Attn Score Recompute
Full-KV (no top-K) analogues of §6 and §7. 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.
9. 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 that overwritesattn_score(target) and readsindex_score(predict) without modifying it.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.
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 treats this buffer as read-only.
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.
For the default SM100 backend, topk_indices_global describes the input
encoding; it is not a performance-tuning switch. When Gather4 is available,
local ids are range-checked and converted to flat ids in registers before the
same optimized gather path. Do not add a separate local_to_global_wrapper
launch only for Indexer Backward; preserve the producer’s id convention and
set topk_indices_global to match it.
On SM90, also prefer preserving the producer’s id convention. The performance difference between local and global ids is generally small, so a separate conversion solely for Indexer Backward is unlikely to help.
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.
Start with the default backend when throughput is the goal. Select sm100_v2 for
its two-term gradient representation, deterministic d_weights, or FP32
output contract, not as an automatic performance-tuned alternative; benchmark
the complete wrapper on the target shape when those properties are required.
- 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_scoreis consumed in place and left holding kernel 1’sgrad_signal, whileindex_scoreis read-only and preserved.sm_scalefolds inside kernel 2 without touching either 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 overwriting
attn_score. 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.
10. 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_k - Constraints — SM90 or SM100+,
H >= 64,ratio >= 1
Limitations
- CuTe DSL requirement — install
nvidia-cutlass-dsl[cu13]>=4.5.0. Sparse Attention Forward emits the few SM100 instructions that are not yet exposed by the public CuTe facade through the MLIR LLVM inline-assembly interface shared by CuTe DSL 4.5 and newer releases. - Architecture support — Sparse Attention Forward supports the mapped
SM100-family capabilities 10.0, 10.3, and 10.7 only.
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+.
- Forward scope — only the 11 supported Prefill instances described above; no SM90, regular H128, FP8 cache, decode, or split-KV forward path.
- Sparse gather path — the public CuTe facade does not yet expose the TMA
gather4operation. Sparse Attention Forward therefore uses a feature-local LLVM inline-assembly bridge with the compiler-owned TMA descriptor; the issued data movement remains hardware TMAgather4. - 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.