Block Sparse Attention (BSA)

View as Markdown

This is an experimental API and subject to change.

Overview

Block Sparse Attention computes non-causal scaled dot-product attention over a block-level sparse pattern. For query token i, let m = floor(i / block_size) be its query-block index and let K_m be the union of the key/value blocks listed for that query block. The operation is

Oi=softmaxj∈Km(QiKjTD)Vj.O_i = \text{softmax}_{j \in K_m} \left(\frac{Q_i K_j^T}{\sqrt{D}}\right)V_j.

The softmax is normalized once over all valid tokens in the selected blocks. block_sizes can shorten individual key/value blocks so that padded tokens do not participate in the softmax.

Unlike NSA Selection Attention, BSA supplies one list of key/value blocks per query block. NSA Selection supplies routing metadata at query-token granularity and is one component of the larger NSA pipeline.

BSA is implemented with Python CuTe DSL/JIT kernels.

Installation

The CuTe DSL runtime is a required dependency, so the base install is enough:

pip install nvidia-cudnn-frontend

The package-wide nvidia-cutlass-dsl[cu13]>=4.6.2 dependency floor applies to BSA. Sage FP8 relies on functionality introduced in CuTe DSL 4.6.1, which the package floor already satisfies. A defensive runtime check remains for environments that force an older DSL.

Forward

The cudnn.torch BSA functions are lazy aliases of the existing cudnn and cudnn.BSA functions, with identical signatures, outputs, and supported configurations.

import torch
from cudnn.torch import (
block_sparse_attention_forward,
block_sparse_attention_backward,
block_sparse_attention_fp8_forward,
)
q = torch.randn(1, 8, 1024, 128, device="cuda", dtype=torch.bfloat16)
k = torch.randn(1, 8, 2048, 128, device="cuda", dtype=torch.bfloat16)
v = torch.randn_like(k)
# For the SM90 blk64 path, this example has 16 Q blocks and 32 KV blocks.
# Every Q block attends to the first four KV blocks.
q2k_block_index = torch.arange(4, device="cuda", dtype=torch.int32)
q2k_block_index = q2k_block_index.view(1, 1, 1, 4).expand(1, 8, 16, 4).contiguous()
block_sizes = torch.full((32,), 64, device="cuda", dtype=torch.int32)
result = block_sparse_attention_forward(
q,
k,
v,
q2k_block_index,
block_sparse_num=4,
block_sizes=block_sizes,
sparse_block_size=64,
)
o, lse = result
# The same values are available as result["o_tensor"] and result["lse_tensor"].

The default layout is BHSD. Pass layout="bshd" for tensors in BSHD layout. The output follows the input layout, while lse is always the FP32 natural-log log-sum-exp with shape (B, H_q, S_q).

Sparse metadata

ArgumentShapeDtypeMeaning
q2k_block_index(B, H_q, N_q, K_max)int32KV-block IDs; only the valid prefix of each last dimension is read
block_sparse_numscalarPython intFixed valid-prefix length for every query block
q2k_block_nums(B, H_q, N_q)int32Optional per-query-block valid-prefix lengths; overrides block_sparse_num
block_sizes(N_kv,) or backend-supported batched formint32Number of valid tokens in each KV block

Here N_q = ceil(S_q / sparse_block_size) and N_kv = ceil(S_kv / sparse_block_size). Except for packed GQA described below, metadata is per query head. The metadata tensors must reside on the same CUDA device as q, k, and v. The value ranges below are a hard caller contract.

Value ranges (caller contract)

Let K_max = q2k_block_index.shape[-1]. Only the active prefix of each q2k_block_index row — the first q2k_block_nums[b, h, m] entries with variable counts, or the first block_sparse_num entries with a fixed count — is consumed. Values in the inactive suffix are ignored.

  • Within each row, active q2k_block_index values must be unique integers in [0, N_kv).
  • For forward with variable counts, every q2k_block_nums value must be in [0, K_max] when allow_empty_block_nums=True, and in [1, K_max] otherwise. Backward variable counts may be in [0, K_max].
  • With fixed counts, block_sparse_num must be in [1, K_max]. The SM100/SM103 blk128 path additionally requires an even value, i.e. an even block_sparse_num in [2, K_max].
  • The block_sizes entry for every physical KV block referenced by an active q2k_block_index value must be in [1, sparse_block_size]. Entries for unreferenced physical KV blocks are ignored. A zero-sized referenced block is not supported; use q2k_block_nums (or the sparse index prefix) to drop the block instead.

Tensor value ranges and per-row uniqueness are not validated at runtime. Violating the contract is unsupported and may produce invalid results or invalid device memory accesses.

On the SM100/SM103 blk128 path, pack_gqa=None automatically packs GQA when the GQA ratio r = H_q / H_kv divides 128. Packed metadata has shape (B, H_kv, ceil(S_q * r / 128), K_max). Pass pack_gqa=False to use the unpacked (B, H_q, ceil(S_q / 128), K_max) contract on every architecture. Explicit pack_gqa=True requires r to divide 128.

When block_sizes=None, each referenced physical KV block is treated as full. Provide block_sizes whenever a referenced final block is only partially valid.

sparse_block_size=None chooses blk64 on SM90/SM120 and blk128 on SM100/SM103. On SM90 or SM120, passing sparse_block_size=128 selects a native KV128 kernel that consumes blk128 metadata directly; it does not expand the metadata or invoke the blk64 kernel. Passing sparse_block_size=64 explicitly selects the SM100/SM103 blk64 CuTe DSL path, whose shape support is narrower. kv_splits is available on SM90 and the explicit Blackwell blk64 path; use_clc applies only to the explicit Blackwell blk64 path.

For SM120 blk128 with fixed block_sparse_num, full physical KV blocks (block_sizes=None), and a KV sequence length divisible by 128, the dispatcher uses an FA4-style native specialization with a dedicated K/V load warp, register-resident Q, and a four-fold-unrolled sparse loop. Most CTAs process 128 Q rows; an underfilled final scheduling wave may use 64 Q rows per CTA while still loading full KV128 blocks and using the parent Q block’s original metadata. This is Q-work scheduling, not KV128-to-KV64 lowering. Variable per-row block counts, explicit block sizes, and partial final KV blocks use the general native blk128 kernel. The FA4-style specialization requires nvidia-cutlass-dsl >= 4.7.0; an older public DSL version is rejected with a version-specific error before the specialized kernel is imported. The package-wide downstream dependency floor is unchanged.

kv_splits=2..256 computes FP32 partial outputs and combines them, with workspace growing linearly in the split count. SM90 accepts an explicit integer split count. The SM100/SM103 blk64 path also accepts kv_splits="auto"; CLC is compatible with split execution. Pass use_clc=True to use persistent CLC scheduling with kv_splits>1, or use_clc=False to use one tile per CTA. The default use_clc=None keeps CLC disabled for split execution because the automatic scheduler policy has not been tuned for that combination. Automatic split selection uses metadata capacity rather than per-row count values. It falls back to a smaller split count when the estimated live workspace does not fit the available CUDA allocator budget; an explicit split count that exceeds that budget raises RuntimeError.

Sage FP8 forward

Sage FP8 is a forward-only blk64 path. Its public wrapper accepts contiguous BF16 Q, K, and V tensors in BHSD layout and performs FP8 quantization internally:

fp8_result = block_sparse_attention_fp8_forward(
q,
k,
v,
q2k_block_index,
block_sparse_num=4,
)
o_fp8 = fp8_result["o_tensor"]

Q, K, and V must be BF16 MHA tensors in contiguous BHSD layout, have matching batch and head counts, and use D=128. The wrapper accepts the same blk64 sparse index metadata used by the regular forward API. It lazily loads the quantizer, creates the E4M3 tensors and FP32 scales needed by the kernel, and does not expose those implementation details as public inputs or outputs.

The result is a one-key TupleDict containing o_tensor, a contiguous BF16 tensor of shape (B, H, S_q, 128). This API does not return LSE and has no backward implementation.

The architecture-specific FP8 contracts are:

  • SM100/SM103 accepts any positive batch and head counts and requires both sequence lengths to be multiples of 64. It uses fixed block_sparse_num with full 64-token KV blocks; q2k_block_nums and block_sizes are not supported. Split-KV is selected internally, and the public FP8 API does not expose kv_splits or use_clc.
  • SM120 accepts any positive batch and head counts, non-aligned Q/KV sequence tails, fixed or per-query-block counts, and block_sizes shaped (N_kv,), (B, N_kv), or (B, H, N_kv). It does not use split-KV.

block_sparse_attention_fp8_forward relies on functionality introduced in CuTe DSL 4.6.1; package-supported installations provide CuTe DSL 4.6.2 or newer. Its internal Q/K/V quantizer is also implemented in CuTe DSL and is loaded lazily, so importing cudnn does not eagerly import it.

Backward

Backward is an explicit API rather than a registered PyTorch autograd operation. It recomputes probabilities from the forward output and LSE:

dout = torch.randn_like(o)
grads = block_sparse_attention_backward(
dout,
q,
k,
v,
o,
lse,
q2k_block_index,
block_sparse_num=4,
block_sizes=block_sizes,
sparse_block_size=64,
)
dq, dk, dv = grads

The result keys are dq_tensor, dk_tensor, and dv_tensor. Optional preallocated dq_tensor, dk_tensor, and dv_tensor arguments are supported. The backward implementation builds a bucketed K-to-Q CSR task layout on the GPU. bucket_size_blocks is an optional tuning override; leaving it unset uses the backend default.

Backward defaults to blk64 on SM90 and blk128 on SM100/SM103. Pass the same explicit sparse_block_size used by forward when selecting the Blackwell blk64 path. SM100/SM103 blk128 backward does not yet consume block_sizes; it therefore requires full physical KV blocks and block_sizes=None.

Current support

Forward

ArchitectureSparse blockPublic input / kernel dtypeQK / V dimensionsAttention
SM9064 or 128 (explicit)FP16, BF16each of 64, 96, 128MHA, GQA, MQA
SM100/SM103128FP16, BF16QK=V=64, 96, or 128MHA, GQA, MQA
SM100/SM10364 (explicit)BF16QK=128, V=128MHA
SM100/SM10364BF16 / FP8 E4M3QK=128, V=128MHA
SM12064FP16, BF16QK=128, V=128MHA, GQA, MQA
SM120128 (explicit)FP16, BF16QK=128, V=128MHA, GQA, MQA
SM12064BF16 / FP8 E4M3QK=128, V=128MHA

SM90 blk64 requires S_q to be a multiple of 64; native blk128 supports positive arbitrary Q/KV lengths, including partial final blocks. Its fixed count may be any positive value. The SM100/SM103 blk128 fixed count must be even and at least two; SM120 and the explicit Blackwell blk64 path accept any positive fixed count. Variable counts use q2k_block_nums. allow_empty_block_nums defaults to False; when it is True, empty rows (q2k_block_nums == 0) produce O = 0 and LSE = -inf. SM90 blk64 selects the empty-row handling as a compile-time specialization, so the default non-empty configuration keeps its branch-free fast path. Its split-KV execution excludes empty rows. Native SM90 blk128 supports empty rows and empty splits with kv_splits=1..256.

Backward

ArchitectureSparse blockDtypeHead dimensionAttention
SM9064 or 128 (explicit)BF16128MHA
SM100/SM10364BF16128MHA
SM100/SM103128BF1664 or 128MHA

Backward is not implemented for SM120. It requires equal QK/V dimensions and does not currently support GQA/MQA.

SM90 native blk128 requires CuTe DSL >= 4.6.2. Q/K/V, O/dO and gradient buffers must have a contiguous head dimension, 16-byte aligned base pointers, and row strides divisible by 16 bytes. BHSD/BSHD layouts and aligned noncompact rows are addressed directly. Sparse metadata can be strided and remains on the GPU; backward makes noncontiguous sparse indices and counts contiguous before building CSR, as in blk64. Forward supports block_sizes of rank 1, 2, or 3; backward supports ranks 1 and 2. Both paths predicate sequence tails even without block_sizes. Forward uses Q128 tiles and shares each KV128 load across two compute warpgroups. Backward uses Q64xKV128 CTAs: both compute warpgroups execute all five GEMMs, splitting KV rows for QK/dP/dK/dV and head dimensions for dQ. Each native Q128/KV128 CSR edge streams two Q64 subtiles without expanding the metadata. Backward accumulates gradients with FP32 atomics and is not deterministic. Benchmark the explicit path for your shapes; blk64 remains the default, and native blk128 is not faster for every workload.

Limitations

The current sparse kernels do not implement causal or local masking, dropout, mask_mod, score_mod, paged KV cache, softcap, or variable-length packed sequences. Regular forward/backward inputs must be rank four, use FP16/BF16 as allowed above, and have a contiguous last (head-dimension) axis. Sage FP8 uses the stricter BF16 BHSD contract described above and quantizes to E4M3 internally. Regular forward outputs are contiguous in the requested BHSD or BSHD layout, including split-KV execution; FP8 output is contiguous BF16 BHSD.

Compilation is lazy. The first call for a new static configuration JIT-compiles the relevant kernel; subsequent calls reuse an in-process cache.

The Torch public surface consists of allocating function wrappers under cudnn.torch, also available under cudnn and cudnn.BSA; there is no separate APIBase class or explicit compile() lifecycle for BSA.

Correctness tests and FP32 references are under test/python/fe_api/bsa.

Acknowledgements

We would like to express our gratitude to huangyitong.hyt@alibaba-inc.com and wenting.swt@alibaba-inc.com for providing testing and optimization feedback throughout the deployment process, which has continuously advanced the BSA kernel toward Speed of Light.

Experimental JAX support

cudnn.jax.block_sparse_attention_forward and cudnn.jax.block_sparse_attention_backward provide explicit forward/backward; cudnn.jax.block_sparse_attention adds first-order reverse-mode differentiation. All three accept JAX arrays, eagerly or under jax.jit, and require no PyTorch. The existing cudnn and BSA Torch APIs are unchanged. The former _jax exports have been removed without compatibility aliases.

The JAX implementation reuses the SM100 blk128 forward, bucketed CSR, backward preprocess, backward, and gradient conversion kernels. Install jax[cuda13] alongside the frontend (CuTeDSL >=4.7, JAX >=0.9.1). The runtime imports no PyTorch; torch parity tests are separate.

The JAX signatures omit Torch-specific launch options and caller-provided output buffers. XLA owns the outputs and workspaces.

Initial contract

PropertySupported
DeviceOne visible CUDA GPU, exactly compute capability 10.0 (SM100)
DataBF16 Q/K/V/O/dO/dQ/dK/dV; FP32 LSE and accumulation
DimensionsMHA with equal head counts, D=64 or 128; positive sequence lengths divisible by 128
LayoutCompact BHSD or BSHD, independently specialized
Sparsity128-token blocks, int32 indices [B,H,Sq/128,C], no block_sizes
CountsFixed even block_sparse_num in [2,C], or runtime int32 q2k_block_nums[B,H,Sq/128]
DifferentiationFirst-order reverse-mode Q/K/V gradients through block_sparse_attention

Variable counts must be in [1,C], or [0,C] with allow_empty_block_nums=True. Only each row’s active prefix is read. Active indices must be unique and in [0,Sk/128). An empty row returns zero output, negative-infinity LSE, and zero query gradient. Unselected K/V receive zero contributions. Metadata values are the caller’s responsibility: the runtime validates shapes, dtypes, and static options, without copying metadata to the host or synchronizing to inspect it. Invalid values are outside the contract.

SM103/SM110/SM120, blk64, FP16/FP8, GQA, different V dimensions, partial blocks, variable sequence lengths, causal masks, dropout, sharding, vmap, JVP, and higher derivatives are unsupported. KDA is outside this change. Static options (including scale and fixed count) specialize compilation; sparse arrays remain runtime operands. Do not mutate saved forward inputs or metadata before backward.

Usage

import jax
import jax.numpy as jnp
from functools import partial
from cudnn.jax import (
block_sparse_attention_forward as forward,
block_sparse_attention_backward as backward,
block_sparse_attention as attention,
)
q = jnp.ones((1, 2, 256, 64), dtype=jnp.bfloat16)
k = jnp.ones((1, 2, 512, 64), dtype=jnp.bfloat16)
v = jnp.ones_like(k)
indices = jnp.broadcast_to(jnp.array([0, 2], jnp.int32), (1, 2, 2, 2))
o, lse = forward(q, k, v, indices, 2) # eager; asynchronous GPU execution
run = jax.jit(partial(forward, block_sparse_num=2))
result = run(q, k, v, indices)
o, lse = result["o_tensor"], result["lse_tensor"]
dq, dk, dv = backward(jnp.ones_like(o), q, k, v, o, lse, indices, 2)
def loss(q, k, v, indices):
return attention(q, k, v, indices, 2).astype(jnp.float32).sum()
grad = jax.jit(jax.grad(loss, argnums=(0, 1, 2)))
dq, dk, dv = grad(q, k, v, indices)
dq.block_until_ready() # needed for host timing, not between GPU operations

Explicit forward returns a JAX pytree with o_tensor, lse_tensor; backward returns dq_tensor, dk_tensor, dv_tensor. Both support tuple unpacking and dictionary access. Explicit helpers do not register autodiff rules; use attention for jax.grad. LSE and metadata are not differentiable public outputs. Backward must receive the exact matching forward Q/K/V/O/LSE, sparse pattern, scale, layout, and empty-row option.

Execution and ownership

Every launch receives XLA’s CUDA stream. Forward is one custom call; backward composes CSR construction, preprocessing, attention backward, and gradient conversion inside one custom call. XLA owns all outputs and workspaces. CSR counts and multi-bucket dK/dV accumulators are initialized for each invocation; preprocessing clears dQ accumulation. No mutable global scratch or public input buffer is donated. Atomic accumulation does not promise bitwise determinism.

TensorSpec.mode presents compact BHSD/BSHD memory in the order expected by each kernel. It does not rewrite DLPack capsules or insert a Python transpose. XLA may still insert layout conversions when surrounding operations use incompatible physical layouts; this is not a universal zero-copy guarantee.

Two cached call builders retain static configuration and kernel objects, never input arrays or pointers. XLA handles compilation and execution; there is no separate plan lifecycle. One custom_vjp registration connects forward and backward. bucket_size_blocks optionally controls backward query buckets; the default reuses the torch path’s heuristic. Multi-device placement and cache portability across architectures require further qualification.

The JAX tests live in test/python/fe_api/jax and are collected by the OSS suite. Torch-free import and gradient checks run in fresh subprocesses. To run the suite in a torch-free container, pass --confcutdir=test/python/fe_api/jax so pytest does not load the torch-based parent conftest.