**English** | [简体中文](./README_zh.md)
Flash-Sparse-Attention is a high-performance trainable sparse attention implementation that combines Flash Attention's memory efficiency with sparse computation for handling extremely long sequences in Transformer models.
> [!NOTE]
> The CuTe backend ([`flash_sparse_attn.ops.cute`](./flash_sparse_attn/ops/cute)) currently delivers the best performance. A Gluon backend targeting performance parity with CuTe is still a work in progress.
# Key Features
> [!NOTE]
> Support for arbitrary mask and bias shapes is available in [this branch](https://github.com/HKUSTDial/flash-sparse-attention/tree/final_mask_version). The current main branch no longer maintains that feature set.
## Supported Features
- Forward and backward passes for dense attention, sparse attention, and gated attention
- Regular batched inputs and varlen inputs
- Causal attention and local window attention
- Arbitrary combinations of Q and KV sequence lengths, with head dimensions up to 256
- Grouped Query Attention and Multi Query Attention
- Sparse softmax threshold control
- Gated attention with gate inputs and configurable gating sparsity
- Flex Local Window Attention with per-head arbitrary window sizes and local ranges
- Split-KV for workload balancing in forward and decode workloads
- Split-QO for workload balancing in backward workloads
- Fused Quant for low-precision computation on hardware without native FP8 support
- Top-k gather KV-cache decode
- Paged Attention
**For complete API documentation, please refer to [here](https://hkustdial.github.io/flash-sparse-attention/)**
## Features We Aim to Support
- KV-Cache Manager
- [TLE](https://github.com/flagos-ai/FlagTree/wiki/TLE) backend support
- [Gluon](https://github.com/triton-lang/triton/tree/main/python/triton/experimental/gluon) backend targeting performance parity with CuTe [WIP]
# Installation
## Requirements
- **Linux**: Ubuntu 22.04 or later
- **Device**: GPU, XPU, NPU, or PPU
- **Python**: 3.9 or later
- **PyTorch**: 2.5.1 or later
- **Triton**: 3.6.0 or later
- **Triton Kernels**: 3.6.0 or later
## Install
Install from PyPI:
```bash
pip install flash-sparse-attn
```
To install from source:
```bash
git clone https://github.com/HKUSTDial/flash-sparse-attention.git
cd flash-sparse-attention
pip install .
```
# Quick Start
## Basic Usage
Below are examples for forward, backward, and decode.
```python
import torch
from flash_sparse_attn.ops.triton.interface import (
flash_sparse_attn_func,
flash_sparse_attn_with_kvcache_func,
)
dtype = torch.bfloat16
device = torch.device("cuda")
batch_size, seqlen, num_heads, num_kv_heads, head_dim = 2, 4096, 32, 8, 128
```
### Forward
Combine flex window, split-KV, fused quant, and sparse softmax for maximum performance.
```python
query = torch.randn(batch_size, seqlen, num_heads, head_dim, dtype=dtype, device=device)
key = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device)
value = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device)
output = flash_sparse_attn_func(
query, key, value,
is_causal=True,
softmax_threshold=1.0,
is_local=True,
is_quant=True,
is_split_kv=True,
)
```
### Backward
Combine flex window, split-QO, fused quant, and low-contribution skipping for maximum backward performance.
```python
query = torch.randn(batch_size, seqlen, num_heads, head_dim, dtype=dtype, device=device, requires_grad=True)
key = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device, requires_grad=True)
value = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device, requires_grad=True)
output = flash_sparse_attn_func(
query, key, value,
is_causal=True,
softmax_threshold=1.0,
is_local=True,
is_quant=True,
is_split_kv=True,
is_split_qo=True,
)
output.sum().backward()
```
### Decode
Combine flex window, split-KV, fused quant, sparse softmax, packed GQA, and Graph for maximum decode performance.
```python
query = torch.randn(batch_size, num_heads, head_dim, dtype=dtype, device=device)
key = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device)
value = torch.randn(batch_size, seqlen, num_kv_heads, head_dim, dtype=dtype, device=device)
def fsa_decode_fn():
return flash_sparse_attn_with_kvcache_func(
query, key, value,
softmax_threshold=1.0,
is_local=True,
is_quant=True,
)
# Warmup
for _ in range(3):
fsa_decode_fn()
torch.cuda.synchronize()
# Capture Graph
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
output = fsa_decode_fn()
# Replay
graph.replay()
```
# Benchmarking
Benchmark scripts are located under [tests](tests/), covering forward, backward, and decoding performance.
## Forward Performance
```bash
# Triton backend
python tests/benchmark_forward.py
# CuTe backend
python tests/benchmark_forward_cute.py
# Gluon backend
# WIP
```
## Backward Performance
```bash
# Triton backend
python tests/benchmark_backward.py
# CuTe backend
python tests/benchmark_backward_cute.py
# Gluon backend
# WIP
```
## Decode Performance
```bash
# Triton backend
python tests/benchmark_decode.py
# CuTe backend
# WIP
# Gluon backend
# WIP
```
# Citation
If you use FSA in your research, please cite:
```bibtex
@misc{shi2026cowindowattentioncausalcoverage,
title={CoWindow Attention: Full Causal Coverage Is a Collective Property},
author={Jingze Shi and Zhangyang Peng and Xianduo Li and Yanlin Qi and Xiaotian Lin and Haoxian Chen and Liangdong Wang and Guang Liu and Yuyu Luo},
year={2026},
eprint={2609.32704},
archivePrefix={arXiv},
primaryClass={cs.AI},
url={https://arxiv.org/abs/2609.32704},
}
@misc{shi2026massallocattentionletattention,
title={MassAlloc Attention: Let Attention Allocate Its Own Compute},
author={Jingze Shi and Zhangyang Peng and Xianduo Li and Yanlin Qi and Xiaotian Lin and Haoxian Chen and Liangdong Wang and Guang Liu and Yuyu Luo},
year={2026},
eprint={2609.32712},
archivePrefix={arXiv},
primaryClass={cs.AI},
url={https://arxiv.org/abs/2609.32712},
}
@misc{shi2025trainabledynamicmasksparse,
title={Trainable Dynamic Mask Sparse Attention},
author={Jingze Shi and Yifan Wu and Bingheng Wu and Yiran Peng and Liangdong Wang and Guang Liu and Yuyu Luo},
year={2025},
eprint={2508.02124},
archivePrefix={arXiv},
primaryClass={cs.AI},
url={https://arxiv.org/abs/2508.02124},
}
```
# Acknowledgments
This project builds upon and integrates several excellent works:
- **[OpenSeek](https://github.com/FlagAI-Open/OpenSeek)** - Kernel development support
- **[Flash-Attention](https://github.com/Dao-AILab/flash-attention)** - Memory-efficient attention computation
- **[NVIDIA CUTLASS](https://github.com/NVIDIA/cutlass)** - High-performance matrix operations library
We thank the open-source community for its contributions to efficient Transformer implementations. 🤗