Standard scaled dot-product attention forms a score matrix whose two long axes are sequence positions. For one attention head,
S = QK^T / sqrt(d)
P = softmax(S)
O = PVthe mathematical definition is compact, but a direct GPU implementation can write the large intermediate matrices S and P to high-bandwidth memory before reading them again. FlashAttention changes that dataflow. It processes blocks of queries, keys, and values in on-chip memory and carries enough row-wise softmax state to combine score tiles exactly.
The result is still ordinary softmax attention, apart from normal floating-point effects. The central change is where intermediate data lives and which values survive between tiles.
A score tile is temporary state
Take a query block Q_i and a key block K_j. Their product creates only a local score tile:
S_ij = Q_i K_j^T / sqrt(d)That tile can be consumed while it is resident in fast on-chip storage. The kernel does not need to preserve every S_ij until all key blocks have been visited.
This matters because the complete score matrix grows quadratically with sequence length. Tiling does not remove the quadratic arithmetic of dense attention. It removes the requirement to materialize the full quadratic score and probability intermediates in off-chip memory.
For a fixed query row, each key tile contributes another segment of logits. The remaining problem is softmax: its denominator depends on all logits in that row, including tiles not processed yet.
Online softmax carries a maximum and a normalizer
A numerically stable softmax subtracts the row maximum before exponentiation. When a row arrives in pieces, the maximum can increase as later pieces appear. FlashAttention therefore carries row-wise state across tiles.
Let m be the maximum seen so far and l the corresponding exponential sum. For a new tile with local maximum m_t, define
m_new = max(m, m_t)
l_new = exp(m - m_new) * l
+ sum(exp(s_t - m_new))where s_t denotes logits from the new tile.
The factor exp(m - m_new) rescales the old contribution into the coordinate system defined by the new maximum. If the new tile contains no larger logit, the factor is one. If it raises the maximum, previous exponential mass is reduced consistently.
This recurrence means the kernel does not need the final maximum in advance. The pair (m, l) is sufficient normalization state for the processed prefix of key blocks.
The output accumulator must be rescaled with the same state
Normalization state alone is not enough because attention also computes a weighted sum of value vectors. The partial output must track the same maximum changes.
For a row-level accumulator a, a simplified update is
a_new = exp(m - m_new) * a
+ sum(exp(s_t - m_new) * V_t)After all key tiles have been processed,
O = a / lThe old accumulator and old normalizer receive the same rescaling factor. That coupling preserves the ratio represented by the accumulated weighted sum.
A kernel can organize these operations in several implementation-specific forms, but the invariant is stable: score tiles may disappear after use because the running maximum, normalization sum, and output accumulator retain the information needed to merge later tiles.
Exact attention does not require a stored probability matrix
The term “exact” separates this mechanism from sparse, low-rank, or approximate attention schemes. FlashAttention reorganizes the evaluation of the same dense softmax expression rather than replacing it with a different attention rule.
That distinction is important for memory analysis. Avoiding a stored N × N probability matrix does not imply that only linear work remains. Every dense query-key pair still participates in the attention computation. The saving comes from reducing reads and writes of large intermediates between high-bandwidth memory and on-chip storage.
Hardware details affect the practical tile dimensions and execution schedule. Shared memory or SRAM capacity, register pressure, head dimension, data type, masking, and accelerator architecture all influence the chosen kernel. Those are implementation constraints around the dataflow, not changes to the softmax identity.
Causal masking fits inside the tiled computation
Autoregressive attention excludes future key positions. In a tiled kernel, the causal boundary intersects some score tiles.
Tiles entirely beyond the permitted region can be skipped or treated as masked. Tiles crossing the diagonal need element-level masking before their logits contribute to the online maximum, normalizer, or output accumulator.
The online recurrence remains valid because masked positions contribute no probability mass. Causality changes which logits participate; it does not require the global score matrix to exist.
Other masking forms can also be incorporated when the kernel can apply them consistently to each score tile. Their exact cost depends on the implementation and mask structure.
Backward computation can trade recomputation for stored intermediates
A conventional backward pass might retain a large probability matrix from the forward pass. FlashAttention-style backward computation instead uses compact saved state and recomputes score or probability information in tiles when needed.
That is a deliberate memory-compute trade. Recomputing local products adds arithmetic, but it avoids storing and later reading a quadratic intermediate. On accelerators where memory traffic is a major cost, additional arithmetic can be preferable to additional high-bandwidth-memory transfers.
The exact saved tensors and recomputation schedule vary across FlashAttention versions and kernels. The architectural point is narrower: the forward dataflow exposes enough compact row state to make a non-materialized attention path possible.
IO behavior is the mechanism
FlashAttention is often summarized as a faster attention kernel, but speed is an outcome rather than its defining semantic property. Its durable mechanism is IO-aware tiling.
Queries, keys, and values are loaded in blocks. Score tiles are formed and consumed near the compute units. Online softmax state merges those tiles without requiring a complete score matrix. The weighted output is accumulated under the same normalization state and eventually written as the attention result.
This preserves dense softmax attention while changing the lifetime and location of its intermediates. The quadratic interaction remains; the quadratic off-chip materialization does not.