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

相关

  • 项目
  • 项目
  • 项目
  • 项目