Transformerの計算コストで最も問題になるのがSelf-Attentionです。系列長 $N$ に対して $O(N^2)$ のメモリと計算量が必要で、長い系列を扱う際のボトルネックになっています。系列長を4096から8192に倍にすると、Attentionのメモリ使用量は4倍になります。
しかし、この問題の本質は計算量だけではありません。実は、標準的なAttentionの実装が遅い最大の原因は、GPUのメモリ階層を無視していることにあります。
2022年にDao et al.が提案したFlashAttentionは、Attentionの計算結果を一切近似せず(exact attention)、GPUのメモリアクセスパターンを最適化することで、2〜4倍の高速化とメモリ使用量のO(N)への削減を同時に達成しました。
FlashAttentionを理解することは、以下のような場面で直接役立ちます。
- 長文脈LLMの理解: GPT-4(128K)やClaude(200K)が長い文脈を処理できる技術的基盤がFlashAttentionです
- 学習の効率化: PyTorchの
torch.nn.functional.scaled_dot_product_attentionはデフォルトでFlashAttentionを使用しており、現代のTransformer学習の標準です - カスタムAttentionの設計: Sliding Window AttentionやSparse Attentionなどの派生手法を理解するための基盤です
本記事の内容
- GPUのメモリ階層(HBM vs SRAM)と演算強度
- 標準Attentionのメモリアクセスパターンの問題点
- タイリング(ブロック分割)によるメモリアクセスの最適化
- Online Softmax(数値安定化とストリーミング処理)
- FlashAttentionのアルゴリズム詳細
- Pythonでのシミュレーション実装
前提知識
この記事を読む前に、以下の記事を読んでおくと理解が深まります。
GPUのメモリ階層
HBMとSRAM — 速度と容量のトレードオフ
FlashAttentionを理解するために、まずGPUの内部構造を簡単に見ましょう。GPUには大きく分けて2種類のメモリがあります。
HBM(High Bandwidth Memory): GPUの「メインメモリ」です。NVIDIA A100の場合、容量は40GBまたは80GBで、帯域は2TB/s。モデルの重みや中間テンソルはここに置かれます。
SRAM(Static RAM): 各ストリーミングマルチプロセッサ(SM)に内蔵されたオンチップメモリです。A100の場合、全SMで合計約20MB、帯域は19TB/s。容量は小さいですがHBMの約10倍高速です。
この構造はCPUのキャッシュ階層と類似しています。CPUでL1/L2キャッシュ(小さいが高速)とDRAM(大きいが低速)があるように、GPUにもSRAM(小さいが高速)とHBM(大きいが低速)の階層があります。

出典: Dao et al., “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”, NeurIPS 2022, Fig.1
原論文の Figure 1 は本記事の全体像を1枚に凝縮しています。左のピラミッドがいま説明したメモリ階層(SRAM 19TB/s・20MB / HBM 1.5TB/s・40GB)、中央が本記事の主役であるタイル化ループ——K・V を外側ループ、Q を内側ループでブロック単位に SRAM へコピーし、$N \times N$ の注意行列を HBM に実体化せずにブロック計算する流れ——です。右のバーグラフは GPT-2 での効果で、PyTorch 標準実装が Matmul・Dropout・Softmax・Mask と別々のカーネルで HBM を往復するのに対し、FlashAttention は融合カーネル1本で約7.6倍高速になっています。
| メモリ | 容量(A100) | 帯域 | レイテンシ |
|---|---|---|---|
| SRAM | ~20MB | ~19 TB/s | ~数ns |
| HBM | 40/80GB | ~2 TB/s | ~数百ns |
| 比率 | 1 : 2000〜4000 | ~10 : 1 | ~1 : 100 |
演算強度(Arithmetic Intensity)
操作の性能がHBMの帯域に律速されるか、演算ユニットに律速されるかは、演算強度(arithmetic intensity)で判断できます。
$$ \text{演算強度} = \frac{\text{FLOPs}}{\text{メモリアクセスバイト数}} $$
A100のFP16演算性能は312 TFLOPS、HBM帯域は2 TB/sなので、演算強度が156以上であれば計算バウンド、それ以下ならメモリバウンドです。
行列積 $\bm{C} = \bm{A}\bm{B}$($\bm{A} \in \mathbb{R}^{M \times K}, \bm{B} \in \mathbb{R}^{K \times N}$)の演算強度は $O(\min(M, N, K))$ に比例し、行列が大きければ計算バウンドになります。
一方、要素ごとの操作(softmax、dropout、マスク適用など)は演算強度が $O(1)$ で、必ずメモリバウンドです。これが標準Attentionの非効率の根本原因です。
では次に、標準Attentionがなぜ非効率なのかを具体的に見ていきましょう。
標準Attentionの非効率
計算の流れとメモリアクセス
標準的なScaled Dot-Product Attentionの計算は以下の4ステップです。$\bm{Q}, \bm{K}, \bm{V} \in \mathbb{R}^{N \times d}$ とします。
ステップ1: $\bm{S} = \bm{Q}\bm{K}^T \in \mathbb{R}^{N \times N}$(行列積)
ステップ2: $\bm{P} = \text{softmax}(\bm{S} / \sqrt{d}) \in \mathbb{R}^{N \times N}$(ソフトマックス)
ステップ3: $\bm{P}$ にdropoutを適用(学習時のみ)
ステップ4: $\bm{O} = \bm{P}\bm{V} \in \mathbb{R}^{N \times d}$(行列積)
問題は、ステップ1と2の間で $N \times N$ の行列 $\bm{S}$ をHBMに書き出し、ステップ2でHBMから読み戻す必要があることです。同様に、$\bm{P}$ もHBMに書き出されます。
系列長 $N = 4096$、$d = 128$ の場合:
$$ \bm{S}, \bm{P} \text{ のサイズ} = N^2 \times 2 \text{ bytes (FP16)} = 4096^2 \times 2 = 32 \text{ MB(各行列)} $$
$\bm{Q}, \bm{K}, \bm{V}$ のサイズは $N \times d \times 2 = 1$ MBなので、中間行列 $\bm{S}, \bm{P}$ のメモリアクセスが圧倒的に支配的です。
HBMアクセス量の分析
標準AttentionのHBMアクセス量を数えましょう。
| 操作 | HBM読み出し | HBM書き込み |
|---|---|---|
| $\bm{S} = \bm{QK}^T$ | $\bm{Q}, \bm{K}$: $2Nd$ | $\bm{S}$: $N^2$ |
| softmax($\bm{S}$) | $\bm{S}$: $N^2$ | $\bm{P}$: $N^2$ |
| dropout($\bm{P}$) | $\bm{P}$: $N^2$ | $\bm{P}$: $N^2$ |
| $\bm{O} = \bm{PV}$ | $\bm{P}, \bm{V}$: $N^2 + Nd$ | $\bm{O}$: $Nd$ |
合計のHBMアクセス量は $\Theta(N^2 + Nd)$ です。$N \gg d$ の場合(長い系列)、$N^2$ の項が支配的になります。
FlashAttentionのアイデアは、この $N^2$ の中間行列をHBMに書き出すことなく、SRAM上でタイルごとに計算を完結させることです。
タイリングとOnline Softmax
タイリング(ブロック分割)
FlashAttentionは、$\bm{Q}, \bm{K}, \bm{V}$ を小さなブロックに分割し、各ブロックをSRAMに載せて計算を行います。
$\bm{Q}$ を行方向に $T_r$ 個のブロック、$\bm{K}, \bm{V}$ を行方向に $T_c$ 個のブロックに分割します:
$$ \bm{Q} = \begin{bmatrix} \bm{Q}_1 \\ \bm{Q}_2 \\ \vdots \\ \bm{Q}_{T_r} \end{bmatrix}, \quad \bm{K} = \begin{bmatrix} \bm{K}_1 \\ \bm{K}_2 \\ \vdots \\ \bm{K}_{T_c} \end{bmatrix}, \quad \bm{V} = \begin{bmatrix} \bm{V}_1 \\ \bm{V}_2 \\ \vdots \\ \bm{V}_{T_c} \end{bmatrix} $$
ブロックサイズ $B_r, B_c$ はSRAMの容量に合わせて決定します。各ブロック $\bm{Q}_i \in \mathbb{R}^{B_r \times d}$、$\bm{K}_j, \bm{V}_j \in \mathbb{R}^{B_c \times d}$ がSRAMに収まる条件は:
$$ (B_r + 2B_c) \times d \times 2 \leq M_{\text{SRAM}} $$
A100のSRAMは約192KB/SM(合計約20MB)なので、$d=128$ の場合、$B_r = B_c = 256$ 程度が可能です。
Online Softmaxの必要性
タイリングの問題は、softmaxが行全体のデータに依存することです。
$$ \text{softmax}(\bm{s})_i = \frac{e^{s_i}}{\sum_{j=1}^{N} e^{s_j}} $$
分母の $\sum_{j=1}^{N} e^{s_j}$ を計算するには、行全体の値が必要です。ブロックごとに計算すると、現在のブロックだけでは正しいsoftmaxが計算できません。
Safe Softmaxの復習
まず、数値安定性のための標準テクニックであるsafe softmaxを復習しましょう。$e^{s_i}$ は $s_i$ が大きいとオーバーフローするため、行の最大値 $m = \max_j s_j$ を引きます:
$$ \text{softmax}(\bm{s})_i = \frac{e^{s_i – m}}{\sum_{j=1}^{N} e^{s_j – m}} $$
これは $m$ が同じ値なら分子・分母でキャンセルするので、数学的に等価です。
Online Softmaxのアルゴリズム
Online Softmax(Milakov & Gimelshein, 2018)は、safe softmaxをストリーミングで計算するアルゴリズムです。データを1パスで処理しながら、最大値と指数和を逐次更新します。
ブロック $j$ を処理する前の状態を $(m^{(j-1)}, \ell^{(j-1)})$ とします。ここで $m$ は「これまでの最大値」、$\ell$ は「これまでの指数和(最大値で正規化済み)」です。
ブロック $j$ のスコア $\bm{s}_j$ を観測したとき:
$$ m^{(j)} = \max(m^{(j-1)}, \max(\bm{s}_j)) $$
指数和の更新では、最大値が変わったことを補正します:
$$ \ell^{(j)} = \ell^{(j-1)} \cdot e^{m^{(j-1)} – m^{(j)}} + \sum_k e^{s_{j,k} – m^{(j)}} $$
$e^{m^{(j-1)} – m^{(j)}}$ の項が補正因子です。新しい最大値 $m^{(j)}$ が前の最大値 $m^{(j-1)}$ より大きい場合、過去の指数和は過大評価されているので、この因子で縮小します。
FlashAttentionにおける出力の逐次更新
FlashAttentionでは、softmaxの計算だけでなく、出力 $\bm{O}$ も逐次更新します。ブロック $j$ を処理した後の出力を $\bm{O}^{(j)}$ とすると:
$$ \bm{O}^{(j)} = \frac{\ell^{(j-1)}}{\ell^{(j)}} \cdot e^{m^{(j-1)} – m^{(j)}} \cdot \bm{O}^{(j-1)} + \frac{1}{\ell^{(j)}} \cdot e^{\bm{s}_j – m^{(j)}} \cdot \bm{V}_j $$
この式の意味を理解しましょう。第1項は過去の出力を「新しい正規化定数 $\ell^{(j)}$ と新しい最大値 $m^{(j)}$ で再正規化」したもの、第2項は「現在のブロックの寄与」です。
全てのブロックを処理し終えると、$\bm{O}^{(T_c)}$ が正確なAttention出力と一致します。この逐次更新のおかげで、$N \times N$ の中間行列を保存する必要がなくなります。
アルゴリズムの全体像が見えてきました。次に、FlashAttentionの完全なアルゴリズムを整理しましょう。
FlashAttentionのアルゴリズム
擬似コード
FlashAttentionの前方パスのアルゴリズムを整理します。
入力: Q, K, V ∈ R^{N×d}(HBM上)
出力: O ∈ R^{N×d}(HBM上)
1. ブロックサイズ B_r, B_c を SRAM容量に基づいて設定
2. Q を T_r = ⌈N/B_r⌉ ブロックに分割: Q_1, ..., Q_{T_r}
3. K, V を T_c = ⌈N/B_c⌉ ブロックに分割: K_1, ..., K_{T_c} と V_1, ..., V_{T_c}
4. O = 0, ℓ = 0, m = -∞ を初期化(HBM上)
5. FOR j = 1 to T_c: # 外側ループ: K, V のブロック
K_j, V_j を HBM → SRAM にロード
FOR i = 1 to T_r: # 内側ループ: Q のブロック
Q_i, O_i, ℓ_i, m_i を HBM → SRAM にロード
S_ij = Q_i K_j^T ∈ R^{B_r × B_c} # SRAM上で行列積
m_ij = rowmax(S_ij)
m_i^{new} = max(m_i, m_ij)
P_ij = exp(S_ij - m_i^{new}) # SRAM上でsoftmax(部分)
ℓ_i^{new} = ℓ_i · exp(m_i - m_i^{new}) + rowsum(P_ij)
O_i = (ℓ_i · exp(m_i - m_i^{new}) · O_i + P_ij V_j) / ℓ_i^{new}
m_i = m_i^{new}, ℓ_i = ℓ_i^{new}
O_i, ℓ_i, m_i を SRAM → HBM に書き戻し
END FOR
END FOR
6. RETURN O
HBMアクセス量の分析
FlashAttentionのHBMアクセス量を計算しましょう。
外側ループの各イテレーション(ブロック $j$)で: – $\bm{K}_j, \bm{V}_j$ を読み出し: $2 B_c d$(1回だけ) – 内側ループ $T_r$ 回で $\bm{Q}_i, \bm{O}_i$ を読み書き: $T_r \times 4 B_r d$
合計(全ブロック):
$$ \text{HBMアクセス} = T_c \times (2B_c d + T_r \times 4B_r d) = 2Nd + 4N^2 d \frac{T_c}{N/B_r} $$
$B_r = B_c = B$ として整理すると:
$$ \text{HBMアクセス} = O\left(\frac{N^2 d^2}{M}\right) $$
ここで $M$ はSRAMの容量です。標準Attentionの $\Theta(N^2 + Nd)$ と比較すると、$M$ が十分大きい場合に大幅な削減が実現できます。
A100の場合、$M = 192$ KB、$N = 4096$、$d = 128$ で計算すると、HBMアクセスは標準Attentionの約5〜10分の1に削減されます。
メモリ使用量
標準Attentionでは $\bm{S}, \bm{P} \in \mathbb{R}^{N \times N}$ を保存する必要があり、メモリ使用量は $O(N^2)$ です。
FlashAttentionでは、$\bm{S}_{ij}, \bm{P}_{ij} \in \mathbb{R}^{B_r \times B_c}$ はSRAM上で計算され、HBMには書き出されません。追加で保存するのは $\bm{O}, \ell, m$ のみで、メモリ使用量は $O(N)$ です。
この $O(N^2) \to O(N)$ の削減により、FlashAttentionは系列長を大幅に伸ばすことが可能になりました。$N = 16384$ でも実用的にAttentionが計算でき、これが現代の長文脈LLMを支える技術的基盤です。
次に、FlashAttentionの核心であるOnline Softmaxとブロック出力更新をPythonで実装して、正しさを確認しましょう。
Pythonによる実装
標準Attentionの実装
まず、比較のために標準的なAttentionを実装します。
import numpy as np
import matplotlib.pyplot as plt
np.random.seed(42)
def standard_attention(Q, K, V):
"""標準的なScaled Dot-Product Attention。"""
N, d = Q.shape
scale = 1.0 / np.sqrt(d)
# S = QK^T / sqrt(d)
S = Q @ K.T * scale # (N, N) — この行列がメモリのボトルネック
# 数値安定化 + softmax
S_max = np.max(S, axis=1, keepdims=True)
P = np.exp(S - S_max)
P = P / np.sum(P, axis=1, keepdims=True)
# O = PV
O = P @ V
return O, S, P # S, P は中間行列(メモリ O(N^2))
FlashAttention風の実装
次に、FlashAttentionのブロック処理とOnline Softmaxを実装します。
def flash_attention(Q, K, V, block_size=64):
"""FlashAttention風のブロック処理 + Online Softmax。"""
N, d = Q.shape
scale = 1.0 / np.sqrt(d)
# 出力と統計量の初期化
O = np.zeros((N, d))
ell = np.zeros((N, 1)) # softmax分母(指数和)
m = np.full((N, 1), -np.inf) # 行ごとの最大値
# ブロック数
T_c = (N + block_size - 1) // block_size
T_r = (N + block_size - 1) // block_size
for j in range(T_c):
# K, V のブロックを取得(SRAM にロードに相当)
kj_start = j * block_size
kj_end = min((j + 1) * block_size, N)
K_j = K[kj_start:kj_end] # (B_c, d)
V_j = V[kj_start:kj_end] # (B_c, d)
for i in range(T_r):
# Q のブロックを取得
qi_start = i * block_size
qi_end = min((i + 1) * block_size, N)
Q_i = Q[qi_start:qi_end] # (B_r, d)
O_i = O[qi_start:qi_end] # (B_r, d)
ell_i = ell[qi_start:qi_end] # (B_r, 1)
m_i = m[qi_start:qi_end] # (B_r, 1)
# ブロック間のAttentionスコア
S_ij = Q_i @ K_j.T * scale # (B_r, B_c) — SRAM上で計算
# 現在ブロックの行最大値
m_ij = np.max(S_ij, axis=1, keepdims=True) # (B_r, 1)
# 新しい最大値
m_new = np.maximum(m_i, m_ij) # (B_r, 1)
# softmax の部分計算(数値安定化済み)
P_ij = np.exp(S_ij - m_new) # (B_r, B_c)
# 指数和の更新
ell_new = ell_i * np.exp(m_i - m_new) + np.sum(P_ij, axis=1, keepdims=True)
# 出力の逐次更新
correction = ell_i * np.exp(m_i - m_new)
O_i = (correction * O_i + P_ij @ V_j) / ell_new
# 状態の更新
O[qi_start:qi_end] = O_i
ell[qi_start:qi_end] = ell_new
m[qi_start:qi_end] = m_new
return O # 中間行列を保存しない — メモリ O(N)
正しさの検証と性能比較
FlashAttention風実装の出力が標準Attentionと一致することを確認し、メモリ使用量を比較します。
# テストデータ
N = 512 # 系列長
d = 64 # ヘッド次元
Q = np.random.randn(N, d).astype(np.float64) * 0.1
K = np.random.randn(N, d).astype(np.float64) * 0.1
V = np.random.randn(N, d).astype(np.float64) * 0.1
# 標準Attention
O_std, S_std, P_std = standard_attention(Q, K, V)
# FlashAttention風(ブロックサイズを変えて比較)
block_sizes = [32, 64, 128, 256]
errors = []
for bs in block_sizes:
O_flash = flash_attention(Q, K, V, block_size=bs)
error = np.max(np.abs(O_std - O_flash))
errors.append(error)
print(f"Block size {bs:4d}: max absolute error = {error:.2e}")
# 可視化
fig, axes = plt.subplots(1, 3, figsize=(18, 5))
# 左: 出力の比較(最初の行)
axes[0].plot(O_std[0, :20], 'o-', label='Standard Attention', markersize=8)
axes[0].plot(flash_attention(Q, K, V, block_size=64)[0, :20], 'x--',
label='FlashAttention', markersize=8)
axes[0].set_xlabel('Dimension')
axes[0].set_ylabel('Output Value')
axes[0].set_title('Output Comparison (first row, first 20 dims)')
axes[0].legend()
axes[0].grid(True, alpha=0.3)
# 中央: ブロックサイズ vs 誤差
axes[1].semilogy(block_sizes, errors, 'o-', color='#00d4ff', markersize=10)
axes[1].set_xlabel('Block Size')
axes[1].set_ylabel('Max Absolute Error')
axes[1].set_title('Numerical Error vs Block Size')
axes[1].grid(True, alpha=0.3)
# 右: メモリ使用量の比較
seq_lengths = [128, 256, 512, 1024, 2048, 4096]
mem_standard = [(n**2 * 2) / (1024**2) for n in seq_lengths] # S, P行列 (MB)
mem_flash = [(n * d * 2 + n * 8) / (1024**2) for n in seq_lengths] # O + stats (MB)
axes[2].plot(seq_lengths, mem_standard, 'o-', label='Standard (S + P matrices)', markersize=8)
axes[2].plot(seq_lengths, mem_flash, 's-', label='FlashAttention (O + stats)', markersize=8)
axes[2].set_xlabel('Sequence Length N')
axes[2].set_ylabel('Additional Memory (MB)')
axes[2].set_title('Memory Usage Comparison')
axes[2].set_yscale('log')
axes[2].legend()
axes[2].grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig('flash_attention_comparison.png', dpi=150, bbox_inches='tight')
plt.show()
左のグラフでは、標準AttentionとFlashAttention風実装の出力が完全に一致していることが確認できます(点とバツが重なっている)。中央のグラフでは、どのブロックサイズでも誤差が $10^{-14}$ 以下(浮動小数点の精度限界)であり、FlashAttentionがexact(厳密)であることが確認できます。
右のグラフが最も重要です。系列長が増加すると、標準Attentionのメモリ使用量は $O(N^2)$ で急速に増大し、$N = 4096$ では32MBの追加メモリが必要です。一方FlashAttentionは $O(N)$ で増加し、同じ系列長でもわずかなメモリしか消費しません。この差は系列長が長くなるほど顕著になり、$N = 32768$ 以上の長文脈処理を実用的にしています。
FlashAttention-2 と FlashAttention-3
FlashAttention-2 の改善
FlashAttention-2(Dao, 2023)は、FlashAttention-1の2倍の高速化を達成しました。主な改善点は:
- ループの入れ替え: 外側ループをQブロック、内側ループをK/Vブロックに変更。これにより、出力 $\bm{O}_i$ の書き戻しが外側ループ1回で済みます
- 非行列積演算の削減: rescalingの計算を最小化し、行列積(tensor core)の比率を最大化
- ワープ間の並列化: GPU内のワープ(32スレッド)間でK/Vブロックを分割して並列処理
これらの改善により、A100でのFP16理論演算性能の約70%を達成し、標準Attentionの約5〜7倍の高速化を実現しています。
FlashAttention-3
FlashAttention-3(Shah et al., 2024)は、NVIDIA Hopper(H100)のハードウェア特性を活用した最新版です。
- 非同期パイプライニング: メモリ読み出しと計算を重畳
- FP8対応: H100のFP8 tensor coreを活用した低精度計算
- ブロック量子化: ブロック単位でFP8に量子化してtensor coreの演算を高速化
H100でFP16の理論演算性能の約75%を達成しています。
まとめ
本記事では、FlashAttentionの仕組みを解説しました。
- GPUのメモリ階層(HBM vs SRAM)を理解することが、Attention最適化の出発点
- 標準Attentionの非効率の根本は、$O(N^2)$ の中間行列をHBMに書き出すメモリアクセス
- タイリングにより、計算をSRAMに収まるブロック単位で実行し、HBMアクセスを $O(N^2d^2/M)$ に削減
- Online Softmaxにより、softmaxとAttention出力をストリーミングで逐次更新し、中間行列の保存を不要にする
- メモリ使用量を $O(N^2)$ から $O(N)$ に削減し、出力は厳密に標準Attentionと一致する
FlashAttentionはPyTorch 2.0以降のデフォルトAttention実装に採用されており、現代のTransformer学習・推論の基盤技術です。
次のステップとして、以下の記事も参考にしてください。