Chapter 26
Online Softmax & FlashAttention
Tiled attention with a running maximum and normalizer, and why memory traffic matters.
Attention looks simple until the score matrix appears. A prompt with tokens has query-key scores per head, and storing those scores can dominate memory traffic. FlashAttention keeps exact softmax attention but computes it tile by tile with an online softmax, so the large matrix never has to live in high-bandwidth memory at once.
26.1 Online softmax
For one row of scores , the stable softmax subtracts the row maximum and uses the normalizer
If the row arrives in blocks, keep a running maximum and running normalizer . Suppose the old blocks have and , while the new block has and . The combined maximum is . Rescale both partial sums to that new reference:
The derivation is just multiplying by . For the old terms, ; the new block has the same identity with . This keeps all exponentials near one, just like ordinary stable softmax, but it works without seeing the whole row at once.
The running maximum is what makes the algorithm order-independent. A later block may contain a larger score than every earlier block, so earlier exponentials must be converted to the new origin before they are added. If the later block is smaller, the old state barely changes and the new block is down-weighted by its distance from the old maximum. Either way, the final and are the same quantities a dense softmax row would have computed.
def online_softmax_normalizer(blocks):
"""Return the max and normalizer for one row split into blocks."""
m = -np.inf
ell = 0.0
for scores in blocks:
block_m = np.max(scores)
new_m = max(m, block_m)
ell *= np.exp(m - new_m)
ell += np.exp(block_m - new_m) * np.sum(np.exp(scores - block_m))
m = new_m
return m, ell
26.2 Tiled attention forward
The same rescaling applies to the value-weighted numerator. For a row of attention output, maintain beside and . When a new key-value block arrives, rescale the old numerator, add the new block’s numerator, and return after the final block.
A naive implementation materializes all scores and probabilities:
def softmax(x, axis=-1):
shifted = x - np.max(x, axis=axis, keepdims=True)
weights = np.exp(shifted)
return weights / np.sum(weights, axis=axis, keepdims=True)
def naive_attention(q, k, v, causal=False):
"""Reference attention that materializes the score matrix."""
scores = q @ k.T / np.sqrt(q.shape[-1])
if causal:
pos = np.arange(q.shape[0])
scores = np.where(pos[None, :] <= pos[:, None], scores, -np.inf)
weights = softmax(scores, axis=-1)
return weights @ v, weights
The tiled version streams key-value blocks. It stores only per-query running state plus the current tile, and the tests assert that it equals naive dense attention and naive causal attention for small random arrays.
The numerator update is the part that turns online softmax into attention. Each block forms unnormalized probabilities relative to its own block maximum, multiplies them by that block’s values, and adds the result to the rescaled old numerator. The division by is delayed until all key blocks have contributed. A causal mask fits naturally: scores for future keys in a tile are set to negative infinity before the block maximum and exponentials are computed.
def tiled_attention(q, k, v, block_size, causal=False):
"""Attention forward pass that streams K,V blocks and never stores T by T."""
tokens, value_dim = q.shape[0], v.shape[1]
m = np.full(tokens, -np.inf)
ell = np.zeros(tokens)
numerator = np.zeros((tokens, value_dim))
q_pos = np.arange(tokens)
scale = np.sqrt(q.shape[-1])
for start in range(0, k.shape[0], block_size):
stop = min(start + block_size, k.shape[0])
scores = q @ k[start:stop].T / scale
if causal:
k_pos = np.arange(start, stop)
scores = np.where(k_pos[None, :] <= q_pos[:, None], scores, -np.inf)
block_m = np.max(scores, axis=1)
has_scores = np.isfinite(block_m)
new_m = np.maximum(m, block_m)
exp_scores = np.zeros_like(scores)
exp_scores[has_scores] = np.exp(
scores[has_scores] - block_m[has_scores, None]
)
alpha = np.exp(m - new_m)
beta = np.zeros(tokens)
beta[has_scores] = np.exp(block_m[has_scores] - new_m[has_scores])
numerator *= alpha[:, None]
numerator += beta[:, None] * (exp_scores @ v[start:stop])
ell = alpha * ell + beta * np.sum(exp_scores, axis=1)
m = new_m
return numerator / ell[:, None]
26.3 Memory and IO
Naive attention writes or keeps a score or probability matrix. Online tiled attention keeps , , and an output numerator per query, so the persistent row state scales as rather than . For , the tested accounting has score elements for the naive matrix and online-state entries.
def attention_memory_elements(tokens):
"""Score storage for naive attention and row state for tiled attention."""
return {"naive_scores": tokens * tokens, "online_state": tokens}
The FlashAttention IO argument is about where bytes move, not changing the mathematical operation [dao2022flashattention]. GPU high-bandwidth memory is large but comparatively slow; on-chip SRAM is small but fast. A tiled kernel loads a block of keys and values into SRAM, updates many query rows with online softmax, and writes the final outputs instead of repeatedly writing and rereading the full score and probability matrices from HBM. That is why exact attention can get faster by using less memory traffic.
This is also why a NumPy listing can teach the idea but not the performance. NumPy still creates ordinary arrays and cannot control GPU SRAM. The listing makes the dependency structure visible: each tile is consumed, folded into row state, and discarded. The production kernel fuses those steps so temporary scores remain close to the compute units instead of becoming large global memory tensors. That fusion is the systems lesson: arithmetic is cheap only when the operands arrive at the right place at the right time.
26.4 Backward by recomputation
The forward pass does not save the probabilities that a textbook backward pass would reuse. Instead, FlashAttention stores compact row statistics such as the running maximum and normalizer, then recomputes score tiles during the backward pass [dao2022flashattention]. The recomputed probabilities are exact enough for the same gradients, and each tile immediately contributes to gradients for queries, keys, and values. This trades extra arithmetic for much less saved activation memory, the same kind of trade used by gradient checkpointing.
Conceptually, backward walks over the same tiles as forward. For each tile it reconstructs the local probabilities from , , and , combines them with the upstream gradient on the output, and accumulates local contributions. Once the contribution has been added, the tile can be forgotten again. The algorithm therefore avoids saving the probability matrix in forward and avoids materializing it in backward.
FlashAttention-2 improves the work partitioning and parallelism so more GPU units stay busy [dao2023flashattention2]. FlashAttention-3 targets newer hardware with asynchronous producer- consumer scheduling and low-precision support while keeping the same exact-attention goal [shah2024flashattention3].
|
In practice
|
Modern LLM training and serving usually call a FlashAttention-style kernel whenever the mask, head size, dtype, and hardware are supported. It is still exact softmax attention, so model quality should match a correct dense implementation up to normal floating-point differences. The speedup depends on sequence length and hardware because the win comes from reducing HBM traffic. Unsupported masks or layouts may fall back to other kernels, so numerical tests should compare against a dense reference on tiny shapes. |
26.5 Teach it
The one-sentence version. FlashAttention is exact attention computed as streaming tiles, using an online softmax so the full score matrix is never stored.
An analogy. Do not spread every receipt across the floor. Keep the current maximum, a running total converted to that maximum, and a running weighted basket.
At the board.
-
Write stable softmax with .
-
Split the row into an old part and a block, then rescale both to .
-
Add the same rescaling to the value-weighted numerator.
-
Circle the matrix that disappeared: scores are produced tile by tile, not stored.
Misconceptions to address.
-
"FlashAttention is approximate." It computes exact softmax attention for supported masks.
-
"The trick is only subtracting the max." The trick is updating the max and rescaling old sums.
-
"Backward must save probabilities." It can recompute them tile by tile.
Check for understanding. Why does changing the running maximum require rescaling the old normalizer?
26.6 Exercises
Why does softmax subtract the row maximum before exponentiating, and why does online softmax need to remember that maximum?
Derive (26.2) from the definition of . Then write the matching update for the numerator .
Explain why materialized attention uses score storage while the tiled forward pass uses persistent row state. Compute the two element counts for .
Modify the tiled attention code to support a block-local causal mask. Test it against naive causal attention on random arrays and uneven block sizes.
References
-
[dao2022flashattention] T. Dao et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. 2022. arXiv:2205.14135
-
[dao2023flashattention2] T. Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. 2023. arXiv:2307.08691
-
[shah2024flashattention3] J. Shah et al. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. 2024. arXiv:2407.08608
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762