flash-algo
**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. 🤗