Relaxed-System-Lab/Flash-Sparse-Attention
🚀🚀 Efficient implementations of Native Sparse Attention
Flash‑Sparse‑Attention (FSA)
概要 – Native Sparse Attention (NSA) の Triton ベースのオープンソース実装であり、カーネルループを再構成することでメモリトラフィックと計算オーバーヘッドを劇的に削減します。大規模言語モデル (LLM) のアテンション層のドロップイン置換として機能し、最新の NVIDIA GPU (Ampere, Hopper 等) で動作します。
重要性 – 稀少アテンションは、非常に長いシーケンス(数万トークン)において、フルアテンションの二次コストを管理可能にする手法です。オリジナルの NSA カーネルは、group‑query‑attention (GQA) のヘッドグループサイズが小さい場合(今日の LLM では一般的)、パディングとアトミック加算のオーバーヘッドに悩まされます。FSA は外側/内側のループを入れ替え、作業を3つの専用カーネルに分割することでこれらのボトルネックを回避し、同じ API を使用しながら学習 (prefill) と推論の両方で最大 2‑3 倍の高速化 を実現します。
主な機能
| 機能 | 詳細 |
|---|---|
| 最適化された Triton カーネル | 3つのカーネル(メインのバッチ処理 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 の構築です。
ベンチマーク
リポジトリには2つのスクリプトが含まれています:
scripts/run_unit_test.sh– フォワード/バックワードの正確性をチェックし、生のカーネルレイテンシを測定します。scripts/run_unit_test_sel_attn.sh– selected‑attention 部分(主要なボトルネック)をベンチマークします。
README には2つのパフォーマンス表が示されています:
- カーネルレベル – FSA のレイテンシを 1 とした場合、NSA やフル Flash‑Attention はブロックサイズ/top‑k に応じて 1.7–2.4 倍低速です。
- エンドツーエンド – LLaMA‑2‑70B などの LLM では、学習ステップのレイテンシが約 1.9 秒 (NSA) から約 1.2 秒 (FSA) に短縮され、プリフィルレイテンシも同様に改善されます。
使用すべきケース
- シーケンス長 ≥ 32k トークンの LLM を学習または提供している場合。
- モデルが GQA を使用しており、KV グループあたりのヘッド数が ≤ 8 である場合(最も一般的な構成)。
- NVIDIA Ampere/Hopper GPU を所有しており、Triton をインストールできる場合。
- すでに密なアテンションに Flash‑Attention を使用しており、モデルコードを書き直さずに稀少アテンションの代替手段を求めている場合。
制限事項 / 今後の課題
- 現在は NVIDIA GPU のみをサポートしており、CPU や AMD には対応していません。
- 実装は Q, K, V のヘッド次元が同じ (≤ 256) であることを前提としています。
- NSA と FSA を動的に切り替える「オンラインプロファイリング」モジュールが将来のリリース(2025年9月)で予定されています。
引用
論文で 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 は、長コンテキストを持つ現代の LLM に対して native sparse attention を実用的にする高性能な Triton カーネルライブラリです。既存の PyTorch/transformers コードにドロップインでき、Ampere/Hopper GPU 上で動作し、同じ API と数値的正確性を維持しながら、オリジナルの NSA 実装に対して 2‑3 倍の高速化を実現します。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト