Relaxed-System-Lab/Flash-Sparse-Attention
🚀🚀 Efficient implementations of Native Sparse Attention
Flash‑Sparse‑Attention (FSA)
简介 – 这是一个基于 Triton 的开源 Native Sparse Attention (NSA) 实现,通过重新排列核心循环来大幅减少内存流量和计算开销。它可作为大型语言模型 (LLM) 中注意力层的直接替换模块,并适用于现代 NVIDIA GPU (Ampere, Hopper 等)。
重要性 – 稀疏注意力是一种将完整注意力机制在处理超长序列(数万个 token)时的二次方成本保持在可控范围内的方法。原始的 NSA 核心在 group‑query‑attention (GQA) 头组大小较小(当今 LLM 中的常见情况)时,会受到填充和原子加法开销的影响。FSA 交换了外层/内层循环,将工作拆分为三个专用核心,并避免了这些瓶颈,在保持相同 API 的情况下,在训练 (prefill) 和推理方面均可提供高达 2‑3 倍的加速。
主要功能
| 功能 | 详细信息 |
|---|---|
| 优化的 Triton 核心 | 三个核心 – 主核心 (批次 query‑to‑KV)、归约核心和在线 softmax – 最小化填充数据处理并移除原子操作。 |
| GQA 感知回退 | 对于 GQA 组大小 ≥ 8 的情况,它会自动回退到原始的 NSA 实现,该实现在这种情况下更快。 |
| 广泛的硬件支持 | 在 NVIDIA A100, H100, H200 (SXM 和 PCIe) 上使用 fp16 和 bf16 进行了测试。 |
| 直接替换 API | FlashSparseAttention 镜像了标准的 torch.nn.Module 接口;您只需要提供 cu_seqlens 和输入张量。 |
| 训练与推理 | 适用于仅前向预填充和完整反向传播 (loss‑backward)。 |
| 兼容性 | 建立在 PyTorch ≥ 2.4, Triton ≥ 3.0, HuggingFace transformers ≥ 4.45 以及官方 Flash‑Attention 2.6.3 核心之上。 |
快速开始
# 1️⃣ 安装依赖 (建议使用 Python 3.10+)
pip install -r requirements.txt # 拉取 torch, triton, transformers, datasets, accelerate, flash‑attn
# 2️⃣ 导入并实例化模块
import torch
from fsa.module.fsa import FlashSparseAttention, RopeConfig
fsa = FlashSparseAttention(
hidden_size=4096,
num_q_heads=4,
num_kv_heads=4,
head_dim=128,
kernel_size=32,
kernel_stride=16,
block_size=64,
topk=16,
init_blocks=1,
local_blocks=2,
window_size=512,
rope_config=RopeConfig(
max_position_embeddings=131072,
head_dim=128,
rope_theta=500000,
rope_scaling={
"factor": 8.0,
"high_freq_factor": 4.0,
"low_freq_factor": 1.0,
"original_max_position_embeddings": 8192,
"rope_type": "llama3",
},
),
).cuda().to(torch.bfloat16)
# 3️⃣ 准备 cu_seqlens (累积长度) 和输入
seqlens = torch.LongTensor([65536, 32768]).int().cuda()
cu_seqlens = torch.cat([torch.zeros(1, dtype=torch.int32, device="cuda"),
torch.cumsum(seqlens, dim=0)], dim=0)
x = torch.randn(cu_seqlens[-1], 4096, device="cuda", dtype=torch.bfloat16)
# 4️⃣ 前向 + 反向 (训练) 示例
y = fsa(x, cu_seqlens)
loss = (y * torch.randn_like(y)).sum(-1).mean()
loss.backward()
与普通 transformer 相比,唯一的额外步骤是构建 cu_seqlens,它对变长批次进行了编码。
基准测试
仓库包含两个脚本:
scripts/run_unit_test.sh– 检查前向/反向正确性并测量原始核心延迟。scripts/run_unit_test_sel_attn.sh– 对 selected‑attention 部分(主要瓶颈)进行基准测试。
README 显示了两个性能表格:
- 核心层级 – FSA 的延迟标准化为 1,而 NSA 和完整 Flash‑Attention 根据块大小/top‑k 的不同,速度慢 1.7–2.4 倍。
- 端到端 – 对于 LLaMA‑2‑70B 等 LLM,训练步骤延迟从约 1.9 秒 (NSA) 降至约 1.2 秒 (FSA),预填充延迟也有类似的改善。
使用时机
- 您正在训练或服务序列长度 ≥ 32k token 的 LLM。
- 您的模型使用 GQA,且每个 KV 组的头数 ≤ 8(最常见的配置)。
- 您拥有 NVIDIA Ampere/Hopper GPU 并可以安装 Triton。
- 您已经依赖 Flash‑Attention 进行密集注意力,并希望在不重写模型代码的情况下获得稀疏注意力替代方案。
限制 / 未来工作
- 目前仅支持 NVIDIA GPU;没有 CPU 或 AMD 路径。
- 实现假设 Q, K, V 的头维度相同 (≤ 256)。
- 预计在未来版本 (2025 年 9 月) 发布一个可以在 NSA 和 FSA 之间动态切换的“在线分析”模块。
引用
如果您在论文中使用 FSA,请引用随附的 arXiv 预印本:
@misc{yan2026fsaalternativeefficientimplementation,
title={{FSA}: An Alternative Efficient Implementation of Native Sparse Attention Kernel},
author={Ran Yan and Youhe Jiang and Zhuoming Chen and Haohui Mai and Beidi Chen and Binhang Yuan},
year={2026},
eprint={2508.18224},
archivePrefix={arXiv},
primaryClass={cs.DC},
url={https://arxiv.org/abs/2508.18224},
}
总结
Flash‑Sparse‑Attention 是一个高性能 Triton 核心库,使 native sparse attention 对于具有长上下文的现代 LLM 变得实用。它可直接放入现有的 PyTorch/transformers 代码中,在 Ampere/Hopper GPU 上执行,并在保持相同 API 和数值正确性的同时,比原始 NSA 实现提供 2‑3 倍的加速。
相关
- 项目
- 项目
- 项目
- 项目