Back to archive
#ai#llm#glossary#aigen

FlashAttention

Attention compares many pairs of text fragments. Saving all intermediate results to a computing card's large memory can make data transfers expensive. You want to perform the same calculation while reducing movement and storage of a huge table.

FlashAttention organizes attention computation in blocks in the chip's small, fast memory. It does not write the full matrix of comparisons and weights to the large memory of the GPU, the card used for parallel computation.

Each chunk contributes part of the result, and the algorithm correctly updates the normalization and sum. You cannot simply compute separate chunk averages and average them: the chunks may have different weights.

The mathematical operation of exact attention is preserved, with the usual differences from computer arithmetic. For dense attention, the number of comparisons still grows quadratically. The gain comes from memory and execution organization, rather than removing dependencies between pairs.

Mechanism source: Dao et al., §2–3.2, Algorithm 1.

Mechanism and details

A small work surface processes one tile from a large supply, limiting movement of the entire contents.

For dense attention, the number of operations still grows quadratically with sequence length. The gain comes from computation and memory organization. During training, some quantities are recomputed in the backward pass rather than stored. “Exact” means the operation itself is not approximated; a different order of floating-point operations can produce small numerical differences.

Averaging block results is not enough

Softmax normalizes weights across all allowed keys. Blocks can have very different sums of unnormalized weights. Each therefore needs the appropriate contribution to the result.

This can be done incrementally by maintaining the maximum score mm, the sum l=∑jesj−ml=\sum_j e^{s_j-m}, and the numerator u=∑jesj−mvju=\sum_j e^{s_j-m}v_j. The result is u/lu/l. When a new block raises the maximum to m′m', the previous ll and uu must be rescaled by em−m′e^{m-m'} before adding new terms. The stable normalization mechanism is explained by Milakov and Gimelshein, §3, Algorithm 3.

An original example uses given scores [0,0,ln⁡3,ln⁡3][0,0,\ln 3,\ln 3] and values [0,0,1,1][0,0,1,1]. Before normalization, the weights have proportions [1,1,3,3][1,1,3,3], so the result is 6/8=0,756/8=0{,}75. The two block results are 0 and 1. Their ordinary mean gives the incorrect 0.5.

The demonstration isolates aggregation for one query and scalar values. A full kernel calculates score blocks from Q and K on the GPU; here, given numbers let you verify the calculation by hand. Switching block size is not a benchmark.

Grouped-Query Attention limits the number of K/V heads, while FlashAttention changes how data is transferred and processed. These techniques can be combined. FlashAttention 2 is a later implementation development; the mechanism described here comes from the first paper.

I use AI-generated content as part of my daily learning process.