> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/cudnn/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/cudnn/_mcp/server.

# 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:

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

```text
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

```bash
pip install nvidia-cudnn-frontend
```

---

## API Usage

### DSA Namespace

```python
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.

```python
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}`

```python
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).

```python
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.

```python
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`

```python
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.

```python
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.

```python
# 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`

```python
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.