DeepSeek Sparse Attention (DSA)

View as Markdown

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:

  1. Sparse Attention Forward – sparse Prefill MQA for the supported SM100 H64/H128 shapes.
  2. Sparse Attention Backward – DSA backward for flat MQA tensors on SM90/SM100.
  3. Indexer Forward – CuTe-DSL score kernel (Q @ K^T, ReLU, head reduce, ratio causal mask) that materializes dense scores.
  4. Combined Indexer Forward + Top-K – SM100 compact score generation, Top-K selection, and optional Top-K softmax in one public API call.
  5. Indexer Top-K – SM90+ CuTe-DSL radix top-K kernel with per-row seq_lens.
  6. Sparse Indexer / Attention Score Recompute – sparse (top-K) recompute of indexer and attention scores for training loss.
  7. Dense Indexer / Attention Score Recompute – dense (full-KV) analogues of the above.
  8. Indexer Backward – three-stage pipeline (score-grad, three GEMMs, dtype cast) for sparse top-K score tensors.
  9. Dense Indexer Backward – full-KV counterpart of Indexer Backward.

Architecture

Q, K, W ──┬─► IndexerForward ──► scores ──► IndexerTopK ──► topk_idxs
└─► IndexerForwardTopK ──► topk_idxs, logits, predict
│
v
SparseAttentionForward ──► 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

API Usage

DSA Namespace

from cudnn import DSA
DSA.SparseAttentionBackward
DSA.sparse_attention_backward_wrapper
DSA.SparseAttentionForward
DSA.sparse_attention_forward_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.local_to_global_wrapper
DSA.compactify_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 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 BF16
    • kv: (total_S_kv, D_qk), with the same dtype as q (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 64
    • attn_sink (optional): (H,) FP32
    • topk_length (optional): (total_S_q,) INT32, clamped to [0, K]
    • softmax_scale (optional runtime float): defaults to 1 / sqrt(D_qk)
    • indexer_topk: compile-time prefix length
  • Outputs — stable TupleDict keys:
    • out: (total_S_q, H, 512), with the same dtype as q
    • max_logits: (total_S_q, H) FP32
    • lse: (total_S_q, H) FP32, KV-only and excluding the sink
    • lse_indexer: (total_S_q, H) FP32, or None for indexer_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}
  • Indexer prefix constraint — indexer_topk <= K; when topk_length is supplied and indexer LSE is enabled, every row must have topk_length >= indexer_topk, matching the tile-boundary snapshot contract.
  • Empty rows — out=0, max_logits=-inf, lse=+inf; an enabled lse_indexer is also +inf.
  • Compilation lifecycle — SparseAttentionForward.compile() validates and primes the API object, while the concrete CuTe kernel is JIT-compiled on its first execute()/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 stream must be a cuda.bindings.driver.CUstream belonging to q.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.

import math
result = DSA.sparse_attention_forward_wrapper(
q,
kv,
topk_idxs,
attn_sink=attn_sink,
topk_length=topk_length,
softmax_scale=1.0 / math.sqrt(q.shape[-1]),
indexer_topk=512,
)
out = result["out"]
max_logits = result["max_logits"]
lse = result["lse"]
lse_indexer = result["lse_indexer"]

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/FP16
    • kv: (total_S_kv, D) (K = V; MQA)
    • out, dout: (total_S_q, H, D_v)
    • lse: (total_S_q, H) FP32
    • attn_sink: (H,) FP32
    • topk_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 expect topk_length <= topk_max. Omitting this tensor uses all topk_max slots.

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}
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,
deterministic=True, # optional; SM100 H16/H32/H64/H96/H128
)
dq, dkv, d_sink = result["dq"], result["dkv"], result["d_sink"]

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 by q_scale * sm_scale.
    • q_causal_offsets (optional): CUDA INT32 tensor with one entry per batch/THD segment, on the same device as q.
  • Output — scores: (B, S_q, S_k) FP32.
  • Precision paths
    • SM90 precision="fp8": Q/K use E4M3 and q_scale/k_scale are FP32 descales with one value per token/head. Set return_lse=True (or provide lse_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_size is currently fixed at 32 and qhead_per_kv_head ∈ {32, 64}. THD inputs additionally require cu_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 times qhead_per_kv_head is 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 supports qhead_per_kv_head ∈ {16, 32, 64} and currently requires H_kv == 1; SM100 BF16 dense and combined Top-K paths support qhead_per_kv_head ∈ {32, 64}, as do their MXFP8 paths. All currently require H_kv == 1 (MQA).
result = DSA.indexer_forward_wrapper(
q, k, w, ratio=4, q_causal_offsets=q_causal_offsets,
)
scores = result["scores"]

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.

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"]

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/BF16
    • seq_lens: (batch_size,) INT32 (per-batch effective column count)
  • Outputs — tuple (indices, values) (values is None when return_val=False). Use return_val=False when only the indices are consumed, so no values output buffer is allocated or written.
  • 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"]

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 (optional topk_length). topk_indices are per-batch local KV ids by default; pass topk_indices_global=True when using ids encoded as batch_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 (optional topk_length). topk_indices are per-batch local KV ids by default; pass topk_indices_global=True when using ids encoded as batch_idx * S_k + local_idx.
  • Output — target: (B, S_q, topk) FP32.
  • Note: the wrapper handles the -log2(e) * lse preprocessing 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:

  1. ScoreGradSm90 / ScoreGradSm100 (kernel 1) — in-place score-grad precompute that overwrites attn_score (target) and reads index_score (predict) without modifying it.
  2. IndexerBackwardSm90 / IndexerBackwardSm100 (kernel 2) — three warp-specialised GEMMs produce d_index_q, d_weights, and a dIndexK_f32 accumulator.
  3. Pure-torch dtype cast (kernel 3) converts dIndexK_f32 to 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 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 their lo term is small). lo is ~2^-9 of hi, 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 in hi and the opposite infinity in lo, which sum to NaN rather than being clamped. d_weights is reduced deterministically — d_weights and d_index_q are bitwise run-to-run stable; d_index_k stays in the same fp32-atomic summation-order class as the default backend.
  • Output dtype selects output precision — d_weights and d_index_k accept 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). FP32 d_index_k is zeroed internally (no caller pre-zero contract). d_index_q is always BF16.
  • Envelope (validated by check_support when 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 == 0 with 128 <= topk <= 2048; sm_scale > 0; contiguous same-device tensors; attn_score/index_score/topk_indices shaped (B, S_q, topk). sm_scale folds 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 score sm_scale * S rounds 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 by S) is what differs. The default backend’s multiply preserves subnormals, so reaching it takes an exact sm_scale * S below 2^-150 (~7.0e-46).
  • Scratch / workspace behavior — attn_score is consumed in place and left holding kernel 1’s grad_signal, while index_score is read-only and preserved. sm_scale folds 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 BF16 d_index_k additionally needs a B * S_k * D fp32 accumulator (2 MiB = 2,097,152 bytes at B=1, S_k=4096, D=128, growing with S_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 via torch.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 (stream when given, otherwise torch.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: cudaStreamPerThread is 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 driving IndexerBackward objects 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-batch S_k before the batch offset is applied, in-kernel: ids < 0 or >= S_k contribute 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,
)

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) BF16
    • weights: (B, S_q, H) BF16
    • index_k: (B, S_k, D) BF16
    • attn_score, index_score: (B, S_q, S_k) FP32 raw dense scores
    • attn_l1norm, index_lse: (B, S_q) FP32 denominators
    • q_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
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

  • 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 gather4 operation. Sparse Attention Forward therefore uses a feature-local LLVM inline-assembly bridge with the compiler-owned TMA descriptor; the issued data movement remains hardware TMA gather4.
  • Indexer Forward only supports head_dim = 128. SM90 supports qhead_per_kv_head ∈ {16, 32, 64} with H_kv = 1; SM100 BF16 and MXFP8 support qhead_per_kv_head ∈ {32, 64}. Both the dense and combined Top-K paths require H_kv = 1.
  • Standalone Top-K only up to 2048; top_k > 2048 is 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 requires qhead_per_kv_head ∈ {32, 64}; explicit microbatching cannot be combined with MXFP8, LSE, or explicit per-batch causal offsets.