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
\( 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#
Install the CuTe DSL optional dependencies:
pip install nvidia-cudnn-frontend[cutedsl]
The package-wide nvidia-cutlass-dsl[cu13]>=4.5.0 dependency floor applies to
BSA. The FP16/BF16 APIs continue to work with supported CuTe DSL 4.5 releases.
The Sage FP8 API below performs an additional runtime check and requires
nvidia-cutlass-dsl>=4.6.1; this narrower requirement does not change the
package dependency floor or make importing cudnn require CuTe DSL 4.6.1.
Forward#
import torch
from cudnn import BSA
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 = BSA.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#
Argument |
Shape |
Dtype |
Meaning |
|---|---|---|---|
|
|
|
KV-block IDs; only the valid prefix of each last dimension is read |
|
scalar |
Python |
Fixed valid-prefix length for every query block |
|
|
|
Optional per-query-block valid-prefix lengths; overrides |
|
|
|
Number 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_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. 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.
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 = BSA.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 requires
B=1,Hequal to 4 or 8, and both sequence lengths to be multiples of 64. It uses fixedblock_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 requires CuTe DSL 4.6.1 or newer at call
time. Its internal Q/K/V quantizer is also implemented in CuTe DSL and is
loaded lazily, so importing cudnn still works with the package-wide CuTe DSL
4.5 dependency floor.
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 = BSA.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#
Architecture |
Sparse block |
Public input / kernel dtype |
QK / V dimensions |
Attention |
|---|---|---|---|---|
SM90 |
64 |
FP16, BF16 |
each of 64, 96, 128 |
MHA, GQA, MQA |
SM100/SM103 |
128 |
FP16, BF16 |
QK=V=64, 96, or 128 |
MHA, GQA, MQA |
SM100/SM103 |
64 (explicit) |
BF16 |
QK=128, V=128 |
MHA |
SM100/SM103 |
64 |
BF16 / FP8 E4M3 |
QK=128, V=128 |
MHA (B=1, H=4 or 8) |
SM120 |
64 |
FP16, BF16 |
QK=128, V=128 |
MHA, GQA, MQA |
SM120 |
64 |
BF16 / FP8 E4M3 |
QK=128, V=128 |
MHA |
SM90 currently requires S_q to be a multiple of 64. 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 selects the empty-row handling as a
compile-time specialization, so the default non-empty configuration keeps its
branch-free fast path. Split-KV execution therefore excludes empty rows.
Backward#
Architecture |
Sparse block |
Dtype |
Head dimension |
Attention |
|---|---|---|---|---|
SM90 |
64 |
BF16 |
128 |
MHA |
SM100/SM103 |
64 |
BF16 |
128 |
MHA |
SM100/SM103 |
128 |
BF16 |
64 or 128 |
MHA |
Backward is not implemented for SM120. It requires equal QK/V dimensions and does not currently support GQA/MQA.
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 current public surface consists of allocating function wrappers under
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.