← All Posts
Deep Learning · Transformers· Efficient execution

FlashAttention: Exact Softmax Without the Full Matrix

You do not need to store every attention probability to compute their weighted sum. FlashAttention tiles the computation and maintains a small set of running statistics. Dense softmax attention remains the target function; the memory-access schedule changes.

The expensive intermediate

Naïve attention forms $S=QK^\top/\sqrt{d_k}$, then a probability matrix $P$, then $PV$. For $n$ query and key positions, each score/probability matrix has $n^2$ entries per head. Writing and reading those intermediates in high-bandwidth device memory can be expensive even when the GPU has fast matrix-multiply hardware.

At $n=8192$, one $n\times n$ matrix in two-byte storage occupies 128 MiB. Multiply by heads and batches and the footprint grows quickly. This calculation concerns one intermediate, not the total activation memory of a training step.

Summarize a block of keys for one query

For a block with scores $s_j$ and value vectors $v_j$, retain

$$m=\max_js_j,\qquad\ell=\sum_je^{s_j-m},\qquad u=\sum_je^{s_j-m}v_j.$$

The final output for that block is $u/\ell$. The maximum stabilizes the exponentials, the scalar $\ell$ is the denominator, and vector $u$ is the unnormalized numerator. For many queries, keep one such triple per query row.

Merge two blocks without revisiting their elements

Let blocks A and B have summaries $(m_A,\ell_A,u_A)$ and $(m_B,\ell_B,u_B)$. Put both on a common exponential scale:

$$m=\max(m_A,m_B),\quad a=e^{m_A-m},\quad b=e^{m_B-m},$$
$$\ell=a\ell_A+b\ell_B,\qquad u=au_A+bu_B.$$

This is exact algebra over real numbers. In floating point, summation order changes rounding, so tiled and untiled results need not be bitwise identical. The important claim is that no kernel approximation or removal of allowed token pairs is required.

A merge that avoids large exponentials

Take scores $(1000,1001,999)$ and scalar values $(2,4,8)$. Split the first two entries into block A and the last into block B. Block A has $m_A=1001$, $\ell_A=e^{-1}+1\approx1.3679$, and $u_A=2e^{-1}+4\approx4.7358$. Block B has $(m_B,\ell_B,u_B)=(999,1,8)$.

The merged maximum is 1001. Scale B by $e^{-2}\approx0.1353$. The denominator is about 1.5032 and numerator about 5.8184, giving output about 3.8707. We never evaluate $e^{1000}$. Direct stable softmax over all three entries gives the same answer.

Merge attention one tile at a time
Advance through score/value blocks and inspect the running maximum, denominator, and weighted output. These are computed statistics, not recorded GPU timings.

Extend the merge to matrix tiles

  1. Load a query tile and initialize per-row running statistics.
  2. Load a key/value tile into fast on-chip memory.
  3. Compute its score tile, apply scaling and the relevant mask, and form local statistics.
  4. Merge the local and running statistics with the rescaling equations.
  5. Continue across allowed K/V tiles, then write the normalized query-tile output.

Queries or key tiles completely excluded by a causal or sparse mask can be skipped. Partially allowed tiles need elementwise masking. A tile containing no allowed key for a row contributes zero denominator and numerator; implementations need a convention that avoids subtracting $-\infty$ from itself. The reference lab handles empty blocks explicitly.

Backward trades storage for recomputation

Training gradients need information about probabilities, but retaining the full probability matrix is not the only option. Store compact row statistics such as log-sum-exp and recompute score/probability tiles during backward. This spends arithmetic to reduce activation storage and device-memory traffic.

Dropout adds a reproducibility requirement: recomputation must recover the same dropout decisions used in the forward pass. Random-number state or a deterministic indexing scheme is part of the kernel design, not an optional detail.

Three claims that must stay separate

ResourceWhat changes?
Dense attention arithmeticStill quadratic in sequence length at fixed head width
Stored attention intermediatesThe full $n\times n$ matrices need not be materialized
Device-memory trafficTiling and reuse reduce costly transfers

The full model still stores Q/K/V or recomputes them, residual activations, MLP activations, parameters, and optimizer state as required. “Linear attention memory” in a kernel result does not mean all transformer memory becomes constant or that the model now uses the recurrent linear-attention mechanism.

Map the algorithm onto hardware

The primary source is FlashAttention. FlashAttention-2 and -3 address how to schedule this work efficiently on GPU hardware. The lab verifies tile merging against a direct stable softmax reference.

Try it: Could you merge the blocks in a different order?

The exact mathematical result is the same because each summary represents the same underlying sums. Floating-point rounding can differ. This mergeability is also useful when attention contributions arrive from different devices.