Probabilistic Heuristics & Bayesian Search

FlashAttention: Exact Attention as an IO-Aware Streaming Computation

FlashAttention is not an approximation to attention. Its core idea is to avoid materializing the N by N attention matrix in HBM by computing tiled attention in SRAM and maintaining online softmax statistics.

Positioning

FlashAttention is best understood as an IO-aware implementation of exact Transformer attention. It does not change the attention definition, and it does not remove the pairwise interaction between every query and every key. Its contribution is more specific: it reorganizes the computation so that the large attention score and probability matrices do not have to be written to and read from high-bandwidth memory.

This distinction matters for long-context language models. Standard attention is expensive not only because it performs many dot products, but also because a naive implementation moves large intermediate matrices through global GPU memory. When sequence length grows, memory traffic can become a practical bottleneck even when the arithmetic units are powerful.

The compute units that perform matrix operations live inside the GPU chip. SRAM, the small memory placed close to those compute units on the chip, can exchange information with them much more efficiently and quickly. HBM, by contrast, is a much larger memory system outside the compute core region, and its effective data exchange is substantially slower, often discussed at roughly an order-of-magnitude disadvantage in this kind of memory-hierarchy argument.

The core idea of FlashAttention is therefore not to naively store all matrix information in HBM and repeatedly send it back to the compute units. Instead, it uses SRAM to make the calculation faster and more efficient. The more important point is that this is not an approximation: FlashAttention computes softmax attention exactly, up to floating-point differences, by maintaining the right online normalization statistics.

Problem setting

For a sequence of length N and head dimension d, scaled dot-product attention is

O = softmax ( QKT d ) V.

The naive implementation follows the formula too literally:

Q, K, V
  |
  v
S = QK^T          full N by N score matrix written to HBM
  |
  v
P = softmax(S)   full N by N probability matrix written to HBM
  |
  v
O = PV

Here S and P are both N×N matrices. For N=4096, this means more than sixteen million entries per attention matrix per head before considering batch size, precision, layers, or backward-pass storage. The mathematical object is attention, but the implementation has materialized large intermediate objects that are not needed as final outputs.

Core idea

FlashAttention computes attention in tiles. It loads a block of queries into on-chip memory, streams blocks of keys and values through it, computes temporary score tiles, updates row-wise softmax statistics, accumulates the output numerator, and discards the temporary scores.

Q_tile, K_tile, V_tile
  |
  v
S_tile = Q_tile K_tile^T
  |
  v
update running max, denominator, and output accumulator
  |
  v
discard S_tile
  |
  v
write final O_tile only

The important point is that FlashAttention does not merely split the attention matrix into smaller matrices and store them. Its stronger idea is to avoid storing the full score matrix S and softmax probability matrix P in HBM at all. Temporary score tiles live only long enough to update the online softmax state and output accumulator.

Mathematical structure

For one query row qi, define scores

sij = qiT kj.

The attention output for that query is

oi = ∑j exp(sij) vj ∑j exp(sij) .

Naively, this looks as if all scores sij must be stored before normalization. The online softmax identity shows that this is not necessary. For keys processed so far, maintain a running maximum

mi = maxj∈processed sij,

a running denominator

ℓi = ∑j∈processed exp( sij -mi ),

and an unnormalized output accumulator

ai = ∑j∈processed exp( sij -mi ) vj.

The term mi is the row-wise maximum used for numerical stability. The term ℓi is the softmax denominator after subtracting that maximum. The term ai is the numerator before final division.

When a new key-value block arrives, define its local maximum and the updated maximum:

miblock = maxj∈block sij, minew = max( miold , miblock ).

Then update the denominator and accumulator:

ℓinew = exp( miold - minew ) ℓiold + ∑j∈block exp( sij - minew ), ainew = exp( miold - minew ) aiold + ∑j∈block exp( sij - minew ) vj.

The exponential rescaling term adjusts all previously accumulated quantities when a larger maximum is discovered in the new block. At the end,

oi = aiℓi.

This is the mathematical reason FlashAttention can stream over key-value blocks without storing the full score vector for each query.

Why it can work

FlashAttention works because softmax normalization can be decomposed into row-wise running statistics. The algorithm only needs the current query tile, the current key-value tile, a temporary score tile, and row-wise statistics. It does not need the full attention row in memory at once.

For query and key block sizes Bq and Bk, SRAM stores objects such as

Qtile ∈RBq×d , Ktile ∈RBk×d , Vtile ∈RBk×d.

It also stores the temporary score tile and row-wise statistics:

Stile = Qtile KtileT ∈ RBq×Bk , m∈RBq , ℓ∈RBq , A∈RBq×d.

If N=4096, the full score matrix has 4096×4096 entries. But if Bq=128 and Bk=64, the temporary score tile has only 8192 entries. SRAM cannot hold the full attention matrix, but it can hold a carefully chosen tile and the running statistics needed to make the streamed computation exact.

Takeaway

FlashAttention is an IO-aware reordering of exact attention computation. It keeps the same attention semantics, avoids materializing the full N×N intermediate matrices, and uses online softmax statistics to stream the computation through fast on-chip memory. The right claim is precise: less HBM traffic and exact softmax attention up to floating-point differences, not a removal of the quadratic interaction structure.

References

  • Stanford CME295 Transformers & LLMs Autumn 2025 Lecture 4 - LLM Training.