Trainable fast and memory-efficient sparse attention
| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 2 个月前 | ||
| 7 个月前 | ||
| 2 个月前 | ||
| 9 个月前 | ||
| 28 天前 | ||
| 1 个月前 | ||
| 1 个月前 | ||
| 1 年前 | ||
| 1 年前 | ||
| 1 年前 | ||
| 5 个月前 | ||
| 4 个月前 | ||
| 1 年前 | ||
| 5 个月前 | ||
| 1 年前 | ||
| 1 年前 | ||
| 5 个月前 | ||
| 1 个月前 | ||
| 1 个月前 | ||
| 5 个月前 | ||
| 4 个月前 | ||
| 1 个月前 | ||
| 5 个月前 |
English | 简体中文
Flash-Sparse-Attention 是一个高性能的可训练稀疏注意力实现,将 Flash Attention 的内存效率与稀疏计算能力相结合,用于在 Transformer 模型中处理超长序列。
主要特性
Note
支持任意形状 mask 和 bias 的版本为这个分支,当前主分支不再覆盖这部分功能。
支持的功能
- Dense Attention、Sparse attention 和 Gated attention 的前向与反向传播
- 常规批次输入与变长(varlen)输入
- 因果注意力与局部窗口注意力
- 任意 Q / KV 序列长度组合,以及小于等于 256 的头维度
- 分组查询注意力(Grouped Query Attention)和多查询注意力(Multi-Query Attention)
- 稀疏 softmax 阈值控制
- Gated attention 支持门控输入,以及控制门控稀疏程度
- Flex Local Window Attention 支持逐头(per-head)的任意窗口大小和局部范围
- Split-KV 适用于前向传播和解码的工作负载均衡
- Split-QO 适用于反向传播的工作负载均衡
- Fused Quant 支持非 FP8 原生支持的硬件使用低精度计算
- Top-k gather KV-cache 解码
- 分页注意力(Paged Attention)
完整 API 文档请参考 这里
我们想要支持的功能
安装
依赖
- Linux:Ubuntu 22.04 或更高版本
- 设备:GPU、XPU、NPU 或 PPU
- Python:3.9 或更高版本
- PyTorch:2.5.1 或更高版本
- Triton:3.6.0 或更高版本
- Triton Kernels:3.6.0 或更高版本
安装
直接安装:
pip install flash-sparse-attn
此外,需要安装 triton_kernels:
pip install "triton_kernels @ git+https://github.com/triton-lang/triton.git@v3.6.0#subdirectory=python/triton_kernels"
如果您希望从源码安装(自动包含所有依赖):
git clone https://github.com/flash-algo/flash-sparse-attn.git
cd flash-sparse-attn
pip install .
通过 HuggingFace Kernel 使用
也可以直接从 HuggingFace Kernel 加载 kernel,无需安装本包:
from kernels import get_kernel
fsa = get_kernel("JingzeShi/flash-sparse-attn", version=1, trust_remote_code=True)
# 前向
out = fsa.flash_sparse_attn_func(q, k, v, is_causal=True)
# 反向
out.sum().backward()
# 解码
out = fsa.flash_sparse_attn_with_kvcache_func(q, k_cache, v_cache)
需要先安装 pip install kernels。
快速开始
基本用法
以下是前向、反向和解码的示例。
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
前向
组合 flex window、split-KV、fused quant 和 sparse softmax 以获得最大性能。
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,
)
反向
组合 flex window、split-QO、split-KV、fused quant 和 low-contribution skipping,以获得最佳反向性能。
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()
解码
组合 flex window、split-KV、fused quant、sparse softmax、packed GQA 和 Graph 以获得最大解码性能。
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,
)
# 预热
for _ in range(3):
fsa_decode_fn()
torch.cuda.synchronize()
# 捕获 Graph
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
output = fsa_decode_fn()
# 重放
graph.replay()
性能
以下基准测试涵盖前向、后向和解码工作负载,以FlashAttention作为基线。
NVIDIA GPU
A100
前向传播性能

反向传播性能

解码性能

H20
前向传播性能

反向传播性能

解码性能

RTX PRO 6000
前向传播性能

反向传播性能

解码性能

基准测试
基准测试脚本位于 tests 下,用于评估前向、反向和解码三类场景下的性能。
前向传播性能
python tests/benchmark_forward.py
反向传播性能
python tests/benchmark_backward.py
解码性能
python tests/benchmark_decode.py
引用
如果您在研究中使用 FSA,请引用:
@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},
}
致谢
本项目基于并集成了几个优秀的工作:
- OpenSeek - 内核开发支持
- Flash-Attention - 内存高效的注意力计算
- NVIDIA CUTLASS - 高性能矩阵运算库
我们感谢开源社区对高效 Transformer 实现的贡献. 🤗