gpu-kernels-explained
← /learn · 09

Softmax and FlashAttention

Online softmax keeps a running maximum and rescales; FlashAttention uses it to never write the score matrix to HBM.

Loading the animation…

Concept

Attention computes S=QK⊤S = QK^\top (an N×NN \times N matrix of scores for a sequence of NN tokens), takes a softmax along each row to get PP, and returns O=PVO = PV. Done as three separate kernels, standard attention writes SS to HBM, reads it back for the softmax, writes PP, and reads it again for the product: four passes over N2N^2 numbers. At N=4096N = 4096 and head dimension 64 that is 136 MB for one head, almost all of it the score matrix.

FlashAttention (Dao et al., 2022) never stores SS. It takes a block of BrB_r queries, streams the keys and values past in blocks of BcB_c, and for each pair computes the Br×BcB_r \times B_c score tile on chip, folds it into a running output, and discards it. The animation is FlashAttention-2's loop order (Dao, 2023): one row of tiles per query block. What goes to and from HBM is QQ and OO once, and KK and VV once per query block, between 2.1 MB (if every re-read of KK and VV hits L2, the "K and V from L2" bar) and 34.6 MB (if none does) at the same size.

The obstacle is the softmax: each output row needs ∑jesj\sum_j e^{s_j} over the whole row before any of it can be normalised, and FlashAttention sees the row a tile at a time. The trick that makes it possible is online softmax.

Loading the animation…

Concept

A safe softmax subtracts the row's maximum before exponentiating, so nothing overflows: pi=exi−m/∑jexj−mp_i = e^{x_i - m} / \sum_j e^{x_j - m}. That seems to need two passes (find mm, then sum), and a third to normalise. Milakov and Gimelshein (2018) showed that one pass is enough: keep a running maximum mm and a running sum ℓ\ell of ex−me^{x - m}. When a new block brings a larger maximum m′m', every term already in ℓ\ell was computed against the old mm and is too large by em′−me^{m' - m}, the same factor for all of them, so multiplying ℓ\ell by em−m′e^{m - m'} fixes them all at once. Step through the animation and watch the bars below the scores: they all shrink together when the maximum rises.

FlashAttention keeps exactly these two numbers per query row, and rescales its partial output row OiO_i by the same factor, which is why it can stream KK and VV through without ever seeing a whole row of scores.

Maths

Online softmax is exact. Let mkm_k and ℓk\ell_k be the maximum and the sum after kk blocks. Suppose ℓk=∑j≤k-th blockexj−mk\ell_k = \sum_{j \le k\text{-th block}} e^{x_j - m_k}. After the next block, with mk+1=max⁡(mk,max⁡blockxj)m_{k+1} = \max(m_k, \max_{\text{block}} x_j),

ℓk emk−mk+1+∑blockexj−mk+1=∑earlierexj−mk+mk−mk+1+∑blockexj−mk+1=∑all so farexj−mk+1,\ell_k\, e^{m_k - m_{k+1}} + \sum_{\text{block}} e^{x_j - m_{k+1}} = \sum_{\text{earlier}} e^{x_j - m_k + m_k - m_{k+1}} + \sum_{\text{block}} e^{x_j - m_{k+1}} = \sum_{\text{all so far}} e^{x_j - m_{k+1}} ,

so by induction ℓ\ell ends as ∑jexj−m\sum_j e^{x_j - m} with mm the true maximum, and pi=exi−m/ℓp_i = e^{x_i - m}/\ell is the ordinary softmax, in exact arithmetic. In floating point the two differ only by the rounding of the different order of additions: the animation's last frame shows the largest difference, of order 10−1610^{-16}.

HBM traffic. With ee bytes per value, standard attention moves e (2Nd+4N2+2Nd)e\,(2Nd + 4N^2 + 2Nd): read QQ and KK, write and read SS, write and read PP, read VV, write OO. FlashAttention-2 reads QQ once (eNdeNd), writes OO once (eNdeNd) and reads KK and VV (2eNd2eNd) once for each of the N/BrN/B_r query blocks:

Bflash=e (2Nd+2N2dBr).B_{\text{flash}} = e\,\Bigl(2Nd + \frac{2N^2 d}{B_r}\Bigr) .

The N2N^2 term survives, but divided by Br/dB_r/d and halved, and with L2 reuse it disappears: the floor is 4eNd4eNd, every value read or written once. This is the paper's Θ(N2d2/M)\Theta(N^2 d^2 / M) result with Br∝M/dB_r \propto M/d for on-chip memory MM.

Code

The online softmax step, cut from src/lib/gpu/model.ts:

const mNew = m > bmax ? m : bmax;
const scale = Math.exp(m - mNew);
let s = 0.0;
for (const v of blk) s += Math.exp(v - mNew);
const lPrev = l;
l = l * scale + s;

And the bytes FlashAttention-2 moves per tile (no L2 reuse in moved, every re-read an L2 hit in once):

if (j === 0) {
  moved += eb * br * d;
  once += eb * br * d;
}
moved += 2 * eb * bc * d;
if (i === 0) once += 2 * eb * bc * d;

One head, forward pass (model, BF16, blocks of 128)

Sequence × head dimStandard attentionFlashAttention, no L2 reuseEvery value once (floor)
4096 × 64136 MB34.6 MB2.1 MB
4096 × 128138 MB69.2 MB4.19 MB
8192 × 128545 MB273 MB8.39 MB