Profile
Back to NewsBack
GitHub Trending 5 min
Reader Mode
HKUSTDial/flash-sparse-attention: Trainable fast and memory-efficient sparse attention

HKUSTDial/flash-sparse-attention: Trainable fast and memory-efficient sparse attention

6 hours ago

flash-algo

English | 简体中文

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) 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. 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

Features We Aim to Support

  • KV-Cache Manager
  • TLE backend support
  • 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:

pip install flash-sparse-attn

To install from source:

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.

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.

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.

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.

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, covering forward, backward, and decoding performance.

Forward Performance

# Triton backend
python tests/benchmark_forward.py

CuTe backend

python tests/benchmark_forward_cute.py

Gluon backend

WIP

Backward Performance

# Triton backend
python tests/benchmark_backward.py

CuTe backend

python tests/benchmark_backward_cute.py

Gluon backend

WIP

Decode Performance

# Triton backend
python tests/benchmark_decode.py

CuTe backend

WIP

Gluon backend

WIP

Citation

If you use FSA in your research, please cite:

@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:

We thank the open-source community for its contributions to efficient Transformer implementations. 🤗
Chat with me