FlashAttentionとは? GPUのメモリ階層から高速化の仕組みを解説
記事のまとめ
- FlashAttentionはAttentionを近似せず、計算順序を変えてGPU上のデータ移動を減らす手法である。
- tilingとonline softmaxにより、巨大なAttention行列をHBMへ保存せずに出力を計算する。
- 計算量は通常のAttentionと同じ二次時間だが、追加メモリは系列長に対して線形になる。
通常のAttentionはなぜGPU上で遅いのか
系列長をN、1ヘッドあたりの次元をdとします。通常のAttentionは、QとKからスコア行列Sを作り、行ごとにsoftmaxを適用した確率行列PとVを掛けます。S=dQK⊤,P=softmax(S),O=PV Q、K、VはN×dですが、SとPはN×Nです。素直な実装では、最初の行列積でSをGPUの大容量メモリへ書き出し、softmaxのためにSを読み戻してPを書き、最後の行列積でPをもう一度読みます。Nが長くなるほど、N×Nの中間行列を何度も読み書きする費用が大きくなります。
行列積そのものはGPUのTensor Coreなどで非常に高速に実行できます。一方、演算器へデータが届かなければ、その性能は活用できません。通常のAttentionは大きな中間行列の移動によってメモリ帯域に律速されやすく、FLOPsだけを減らしても実時間が同じ比率で短くなるとは限りません。HBMとSRAM
GPUには性質の異なる複数のメモリがあります。FlashAttentionの理解では、GPU外部の大容量メモリであるHBMと、チップ上にある小容量のSRAMを区別することが重要です。- HBM(High Bandwidth Memory):容量は大きいが、演算器から遠く、データ転送のコストが高い。Q、K、Vや通常のAttentionの中間行列は主にここに置かれる。
- SRAM:各Streaming Multiprocessorに近いオンチップメモリで、HBMより小さいが高速である。CUDAではshared memoryやregisterなどが高速な作業領域として使われる。
SRAMにはN×NのAttention行列全体を置けません。しかし、小さなブロックなら置けます。そこで「必要な部分だけをHBMからSRAMへ運び、SRAM上でできるだけ多くの計算を済ませる」という方針が生まれます。IO complexityという考え方
IO complexityは、計算に必要な加算・乗算の回数ではなく、異なるメモリ階層の間で読み書きするデータ量を評価します。系列長N、ヘッド次元d、SRAMに置ける要素数をMとすると、原論文が示すHBMアクセス量は次のようになります。Standard Attention: Θ(Nd+N2) FlashAttention: Θ(MN2d2) 一般的なヘッド次元ではd²よりMが十分に大きいため、FlashAttentionは通常実装よりHBMアクセスを大幅に減らせます。ここで減るのはAttentionの数学的な計算量ではなく、主にデータ移動量です。FlashAttentionが「同じO(N²d)の演算なのに速い」理由はこの違いにあります。引用:https://arxiv.org/abs/2205.14135
tiling:行列を小さく分ける
tilingとは、Q、K、VをSRAMへ収まる行ブロックに分ける処理です。QのブロックQᵢとK、VのブロックKⱼ、VⱼだけをSRAMへ読み込み、局所的なスコアを計算します。Sij=QiKj⊤/d このSᵢⱼは小さいためSRAM上に保持でき、softmaxとVⱼとの積まで同じカーネル内で処理できます。計算が終わればSᵢⱼを捨て、次のK、Vブロックへ進みます。そのため、完全なN×NのSやPをHBMへ書く必要がありません。
ただし、softmaxの分母は行全体のスコアに依存します。一つのブロックしか見ていない段階では最終的なsoftmaxを決められません。この問題を解くのがonline softmaxです。softmaxをブロック単位で計算する仕組み
数値的に安定したsoftmaxでは、各行の最大値mを引いてから指数を計算します。ある行のスコアx₁,…,xₙについて、分母ℓは次のように表せます。m=kmaxxk,ℓ=k∑exp(xk−m) ブロックを追加すると最大値が変わる可能性があります。FlashAttentionは、それまでの最大値m、指数和ℓ、Vを掛けた出力の分子を行ごとに保持します。新しいブロックの最大値をm̃、新しい指数和をℓ̃とすると、統合後の値は次のように更新できます。mnew=max(m,m) ℓnew=em−mnewℓ+em−mnewℓ 古い最大値を基準にした指数和へeの補正係数を掛ければ、新しい最大値を基準にした値へ変換できます。出力の分子にも同じ補正を行うことで、全ブロックを見終えた時点の結果は、行全体へ一度にsoftmaxを適用した結果と一致します。つまり、各ブロックのsoftmaxを独立に計算して単純に連結しているわけではありません。正規化統計を更新し続けることが正確性の鍵です。FlashAttentionのアルゴリズム
FlashAttentionのforward passは、概念的には次の順序で進みます。- Q、K、VをSRAM容量に合わせたブロックへ分割する。
- KとVのブロックをHBMからSRAMへ読み込む。
- Qの各ブロックと、対応する途中出力O、行最大値m、指数和ℓを読み込む。
- SRAM上で局所スコアQᵢKⱼᵀを計算し、マスクとスケーリングを適用する。
- online softmaxの式でmとℓを更新し、局所確率とVⱼから出力Oを更新する。
- 更新したO、m、ℓだけをHBMへ戻し、次のブロックを処理する。
backward passでもN×NのAttention行列を保存しません。forward passで保存した出力とsoftmaxの正規化統計を利用し、必要な局所スコアと確率をSRAM上で再計算します。再計算によってFLOPsは増えますが、巨大な中間行列をHBMから読み書きするより速くなるという選択です。通常のAttentionとの違い
計算結果と計算量
FlashAttentionは近似Attentionではなく、通常のAttentionを計算順序だけ変えて実装したexact attentionです。forward passの漸近的な演算量はどちらもO(N²d)であり、系列長に対する二次時間そのものは消えません。そのため、FlashAttentionを使えば無制限に長い系列を安価に処理できる、という意味ではありません。FLOPs=O(N2d) メモリ使用量
通常実装はbackward passのためにN×NのAttention確率を保持しやすく、系列長に対して二次の追加メモリが必要です。FlashAttentionは行ごとの統計と出力を保存し、Attention行列を再計算するため、入力と出力を除く追加メモリをO(N)にできます。原論文の実験でもメモリ使用量は系列長に対して線形に増加しています。Extra memory: Standard O(N2),FlashAttention O(N) メリット
- N×NのAttention行列をHBMへ保存しないため、長い系列でメモリ使用量を大きく削減できる。
- HBMとSRAM間のデータ移動が減り、同じAttentionの式をより短時間で計算できる。
- 近似によるモデル品質の低下を導入せず、通常のAttentionを置き換えられる。
- メモリ削減によって、より長いコンテキストや大きいバッチサイズを扱いやすくなる。
- causal mask、dropout、MQA/GQAなど、実用的なTransformer機能に対応する実装が公開されている。
原論文では、最適化済みのベースラインに対してBERT-largeのend-to-end学習を15%高速化し、系列長1KのGPT-2学習を3倍高速化したと報告しています。ただし、速度向上率はGPU、データ型、系列長、ヘッド次元、マスク、他レイヤーの割合によって変わるため、すべての環境で同じ倍率になるわけではありません。引用:https://arxiv.org/abs/2205.14135
デメリット・制約
- 系列長に対する演算量はO(N²d)のままであり、二次時間を線形時間へ変える手法ではない。
- 短い系列や小さい問題では、tilingやカーネル起動の管理コストにより効果が小さい場合がある。
- backward passでは中間行列を保存しない代わりに再計算を行うため、純粋なFLOPsは増える。
- 高性能を得るにはGPU世代、CUDAまたはROCm、データ型、ヘッド次元に適した専用カーネルが必要になる。
- 数学的にはexactでも、浮動小数点の演算順序が変わるため、通常実装とビット単位で同一とは限らない。
公式実装の対応条件もバージョンごとに変わります。現在のFlashAttention-2のCUDA実装は主にAmpere、Ada、Hopper世代、fp16またはbf16、最大256のヘッド次元を対象としています。利用時は論文だけでなく、使用するバージョンのREADMEとテスト条件を確認する必要があります。引用:https://github.com/Dao-AILab/flash-attention
FlashAttention-2との関係
2023年に発表された"FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning"は、FlashAttentionのIO-awareな基本原理を維持しながら、GPU上での仕事の分け方を改善した後続手法です。- softmaxの再スケーリングなど、Tensor Coreを使わない非行列積の演算を減らす。
- batchとheadだけでなく系列長方向にもthread blockを分け、長い系列や小さいバッチでGPUの稼働率を高める。
- warp間の分割をsliced-Kからsliced-Qへ変更し、shared memoryを介した同期と通信を減らす。
FlashAttention-2論文では、A100上で初代FlashAttentionのおよそ2倍、理論ピークの最大73%に達するforward性能が報告されています。これはAttentionの数式を再び変えたのではなく、同じアルゴリズムをGPUのthread blockとwarpへより効率よく割り当てた結果です。その後もHopper向けのFlashAttention-3などが開発されていますが、「演算だけでなくメモリ移動とハードウェア上の仕事分割を設計する」という中心思想は共通しています。引用:https://arxiv.org/abs/2307.08691
まとめ
FlashAttentionの高速化は、Attentionの二次計算を魔法のように消すものではありません。通常実装がN×Nの中間行列をHBMへ何度も読み書きするのに対し、FlashAttentionはtilingによって小さなブロックをSRAM上で処理し、online softmaxによって正規化統計と出力だけを更新します。
その結果、演算量O(N²d)と通常のAttentionの正確な結果を保ちながら、HBMアクセスを減らし、追加メモリをO(N)にできます。FlashAttentionを理解する上で最も重要なのは、「アルゴリズムの速さはFLOPsだけでは決まらず、データをどこへ何回運ぶかにも左右される」という点です。
ご愛読ありがとうございます。