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 倍的加速。

相關

  • 專案
  • 專案
  • 專案
  • 專案