Block Sparse Attention (BSA)
Block Sparse Attention (BSA)
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
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:
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.
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
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_indexvalues must be unique integers in[0, N_kv). - For forward with variable counts, every
q2k_block_numsvalue must be in[0, K_max]whenallow_empty_block_nums=True, and in[1, K_max]otherwise. Backward variable counts may be in[0, K_max]. - With fixed counts,
block_sparse_nummust be in[1, K_max]. The SM100/SM103 blk128 path additionally requires an even value, i.e. an evenblock_sparse_numin[2, K_max]. - The
block_sizesentry for every physical KV block referenced by an activeq2k_block_indexvalue must be in[1, sparse_block_size]. Entries for unreferenced physical KV blocks are ignored. A zero-sized referenced block is not supported; useq2k_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:
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_numwith full 64-token KV blocks;q2k_block_numsandblock_sizesare not supported. Split-KV is selected internally, and the public FP8 API does not exposekv_splitsoruse_clc. - SM120 accepts any positive batch and head counts, non-aligned Q/KV sequence
tails, fixed or per-query-block counts, and
block_sizesshaped(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:
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
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
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
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
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.