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.shselected‑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 倍の高速化を実現します。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト