Built on top of FlashAttention by Tri Dao et al. This fork adds a speculative + sparse attention path on FA-3 (Hopper) so a decoding step can attend to a caller-chosen subset of K/V blocks, then "repair" the result over the blocks it missed — all merged exactly via the FlashAttention online-softmax combine, with no recomputation on the overlap.
The vanilla FA-2 and FA-3 forward/backward kernels are unchanged and continue to work. To coexist with an upstream flash-attn install in the same virtualenv, the packages here are renamed to sparse-flash-attn-2 / sparse-flash-attn-3.
-
Sparse block-table forward (FA-3). A new forward path that attends only to the K/V blocks listed in a
sparse_block_table(1Dint32). Each entry covers ablock_size-token slice ofk/v;block_size ∈ {16, 32, 64, 128, 256}at hdim=128. Implemented inhopper/mainloop_fwd_sm90_tma_gmma_ws.hppviaload_sparse()+mma_sparse(), surfaced asflash_attn_with_sparse_block_table(). -
Two-pass speculative + sparse decoding with online-softmax combine. Run dense FA over a speculative K/V subset (optionally causal), run sparse FA over the missed blocks, merge with the existing
FlashAttnFwdCombinekernel — exact, no recomputation on the overlap. Surfaced asflash_attn_speculative_sparse(). -
Per-block partials.
flash_attn_with_sparse_block_table_partials()returns one(o, lse)per block so callers can choose the subset to combine themselves. -
Validated combine identity.
FA(A ∪ B) == Combine(FA(A), FA(B))is verified numerically in fp64 bytests/test_sparse_combine_identity.py(38 passing cases — no GPU required). -
FA-2 sparse path (
csrc/flash_attn/): the samesparse_block_tablecontract is wired throughcompute_attn_1rowblock_sparse()for Ampere/Hopper. The Python surface is the standard FA-2 API; passsparse_block_tablethroughtorch.ops.sparse_flash_attn_2.fwd.
- NVIDIA Hopper GPU (H100 / H800, SM90) to run the FA-3 sparse path. Ampere works for FA-2.
- CUDA ≥ 12.3 with
nvcconPATH; CUDA 12.8 recommended. - PyTorch built against the same CUDA major version.
ninja,packaging,psutil.
Build from source — there are no prebuilt wheels for the renamed packages.
FA-3 (Hopper, primary):
cd hopper
MAX_JOBS=4 pip install -e . --no-build-isolationFA-2 (optional, Ampere/Hopper fallback):
MAX_JOBS=4 pip install -e . --no-build-isolationCold builds take 20–60 min. Drop MAX_JOBS to 2 if nvcc OOMs. For faster dev iteration, disable head-dims you don't need (FA-3 example):
FLASHATTENTION_DISABLE_HDIM64=TRUE FLASHATTENTION_DISABLE_HDIM192=TRUE \
FLASHATTENTION_DISABLE_HDIM256=TRUE MAX_JOBS=4 pip install -e . --no-build-isolationVerify the FA-3 install registered correctly:
python -c "import torch, sparse_flash_attn_3._C; print([o for o in dir(torch.ops.sparse_flash_attn_3) if not o.startswith('_')])"The printed op list should include fwd, bwd, etc. If you see RuntimeError: operator ... has already been registered, an old flash-attn install is still in the venv and is colliding — uninstall it, or import the renamed package in a fresh process.
See instruction.md for the full build / rename reference.
The new functions live in hopper/flash_attn_interface.py and are importable as flash_attn_interface after cd hopper.
from flash_attn_interface import flash_attn_with_sparse_block_table
# q: (B, Sq, H, D)
# k, v: (B, Sk, Hk, D) — Sk must be a multiple of block_size
# sparse_block_table: (N,) int32 CUDA — absolute block indices into k/v
out, lse = flash_attn_with_sparse_block_table(
q, k, v, sparse_block_table, block_size=128,
)
# out: (B, Sq, H, D), lse: (B, H, Sq)Each entry of sparse_block_table indexes a block of block_size consecutive K/V tokens (physical offset = block_idx * block_size). Every selected block must be fully historical with respect to every Q row — the kernel applies no causal or local mask. Pass an empty table to get out = 0, lse = -inf (a neutral partial for combine).
from flash_attn_interface import flash_attn_speculative_sparse
# k_spec, v_spec: speculative dense K/V subset (must include the causal diagonal)
# k, v: the full K/V that sparse_block_table indexes into
# sparse_block_table: (N,) int32 — the *missed* blocks not covered by k_spec/v_spec
out, lse = flash_attn_speculative_sparse(
q, k_spec, v_spec, k, v, sparse_block_table, causal_spec=True,
)
# out: (B, Sq, H, D) in q's dtype, lse: (B, H, Sq) fp32Internally this runs FA over (q, k_spec, v_spec) with causal_spec masking, runs the sparse pass over the missed blocks (always unmasked), and merges the two (o, lse) partials through the standard FlashAttnFwdCombine kernel. The result is bit-equivalent (up to fp accumulation order) to attending over the union of both block sets.
For callers orchestrating their own combine (e.g. selecting different block subsets per query block), flash_attn_with_sparse_block_table_partials() returns one (o_b, lse_b) per entry in sparse_block_table; feed them into flash_attn_combine yourself.
| Path | Contents |
|---|---|
hopper/ |
FA-3 (CUTLASS 3.x, SM90 TMA+WGMMA+WS). Sparse work lives here. Python entry: hopper/flash_attn_interface.py. |
csrc/flash_attn/ |
FA-2 kernels (CUTLASS 2.x). Sparse path in flash_fwd_kernel.h + compute_attn_1rowblock_sparse. |
sparse_flash_attn_2/ |
FA-2 Python package (renamed from upstream flash_attn/). |
tests/, benchmarks/, examples/, training/ |
Correctness tests, perf benchmarks, example scripts, training reference code. |
# Reference combine identity — CPU, no GPU needed
pytest tests/test_sparse_combine_identity.py -v
# FA-3 sparse path — requires an SM90 GPU
pytest hopper/test_flash_attn.py -k sparse
# FA-2 regression
pytest tests/test_flash_attn.pyIf you use this work, please also cite the upstream FlashAttention papers it is built on:
@misc{wang2026predictreuserepairaccelerating,
title={Predict, Reuse, and Repair: Accelerating Dynamic Sparse Attention for Long-Context LLM Decoding},
author={Tianyu Wang and Gourav Rattihalli and Aditya Dhakal and Junbo Li and Zhiwei Ren and Dejan Milojicic and Longfei Shangguan},
year={2026},
eprint={2606.30389},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2606.30389},
}
@inproceedings{dao2022flashattention,
title={FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness},
author={Dao, Tri and Fu, Daniel Y. and Ermon, Stefano and Rudra, Atri and R{\'e}, Christopher},
booktitle={NeurIPS},
year={2022}
}
@article{dao2023flashattention2,
title={FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning},
author={Dao, Tri},
year={2023}
}
@article{shah2024flashattention3,
title={FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision},
author={Shah, Jay and Bikshandi, Ganesh and Zhang, Ying and Thakkar, Vijay and Ramani, Pradeep and Dao, Tri},
year={2024}
}- Predict, Reuse, and Repair: https://arxiv.org/pdf/2606.30389v1
- FlashAttention paper: https://arxiv.org/abs/2205.14135
- FlashAttention-2 paper: https://tridao.me/publications/flash2/flash2.pdf
- FlashAttention-3 paper: https://tridao.me/publications/flash3/flash3.pdf
BSD-3-Clause, inherited from upstream FlashAttention. See LICENSE.