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 倍的加速。
相關
- 專案
- 專案
- 專案
- 專案