for each query tile:
initialize per-row m, l, u as empty states
for each key/value tile:
compute scaled scores and the global-position mask
for each row with valid keys in this tile:
choose the new maximum
rescale BOTH the old l and old u
add this tile's mass and weighted values
normalize each nonempty row once; handle empty rows explicitly
How Does FlashAttention Preserve Global Softmax Across Tiles?
Derive mergeable maxima, denominators, and unnormalized weighted sums from a tile-local counterexample; check masks and empty rows, then separate backward recomputation, HBM traffic, prefill, and decode.
Cutting an attention matrix into tiles sounds like a memory optimization. Yet the softmax denominator depends on the entire row. A tile cannot know how large a later score will be. How can it contribute now, and why does normalizing each tile separately change the answer? These questions provide a useful entry into FlashAttention.
This article examines the foundations in FlashAttention (2022) and FlashAttention-2 (2023), rather than presenting them as new releases. Starting from one query row, we derive a mergeable state and connect it to masks, training recomputation, and cost. We executed reproducible pure-Python CPU forward checks. We did not run a FlashAttention GPU kernel, model training, or latency benchmarks.
1. Fix the operator before saving memory
Consider one attention head in one batch item. Queries, keys, and values have the three two-dimensional shapes below. The query/key dimension \(d_k\) need not equal the value dimension \(d_v\). Let \(J_i\) contain the keys visible to query row \(i\), and let \(b_{ij}\) be an optional fixed bias independent of the variables differentiated here. Assume finite real logits at allowed positions, and disable dropout initially.
Weights sum to one on every nonempty row; masked positions have zero weight. The scale \(1/\sqrt{d_k}\) does not change with tile size. Switching implementations must also preserve Q/K/V, positional processing, and biases. We change the execution order of the same operator, rather than train a different attention structure.
A materialized implementation writes \(S\) to device memory, applies row-wise softmax to obtain \(P\), and computes \(PV\). With \(N_q=N_k=N\), each intermediate has \(N^2\) entries. Memory must accommodate those entries and transfer them between operations. The FlashAttention question is whether scores can be generated in small tiles, consumed, and discarded while retaining the globally normalized output.
2. Why adding tile-local softmax outputs changes the answer
Use one query and two keys, with one key per tile. Choose logits \((0,\log 3)\) and scalar values \((0,1)\). Global weights are \((1/4,3/4)\), giving output \(3/4\). Each single-element tile has local weight one, so the tile outputs are zero and one. Their sum is one and their average is \(1/2\). Both discard the tiles' different contributions to the global denominator.
This is not merely a small numerical discrepancy. Increasing the logit gap drives the correct output toward one value, while the local outputs still lack information about the other tile's strength. Weighting by tile length does not repair the issue: equal lengths do not imply equal exponential mass.
A tile therefore cannot return only a normalized output. It must supply enough information to establish how much that output should contribute under a global denominator. Numerical stability additionally rules out relying directly on potentially overflowing \(\exp(s_j)\).
3. Keep a maximum, a denominator, and an unnormalized weighted sum
Fix one query row and suppress its row index. For a nonempty allowed-key subset \(A\), define a scalar maximum \(m_A\), scalar normalizer \(\ell_A\), and vector \(u_A\) of length \(d_v\).
The vector \(u_A\) is not the final attention output: it still carries the tile's exponential mass. The maximum defines an exponential coordinate origin, the denominator records the mass in those coordinates, and the weighted sum records which values that mass carries. When \(m_A\) changes, both other quantities must change consistently.
For \(n_A\) finite allowed logits, \(1\le\ell_A\le n_A\): at least one exponential equals one, and all others lie between zero and one. This avoids positive exponential overflow, but does not guarantee freedom from every floating-point error. Tiny terms can underflow, signed values can cause cancellation in the weighted sum, and QK dot products can themselves overflow. Our CPU checks use finite inputs; accumulation precision in a low-precision kernel requires separate investigation.
Online normalizer calculation for softmax (2018) already gave the online maximum and normalizer update. Attention adds a vector weighted sum, transformed in the same coordinates. The figure shows that relationship; the full equations remain in the text.
Original mechanism: split one row’s valid keys into two blocks and retain each maximum, shifted exponential sum, and unnormalized weighted value sum. Merge under a shared maximum, applying each block’s exponential scale to both its normalizer and numerator before summing and normalizing. Do not directly add normalized block outputs. Skip empty blocks; define a policy for an empty row. The rearrangement avoids full attention intermediates in HBM without reducing valid dense-attention pairs. This is not a performance measurement.
Suppose \(A\) and \(B\) are disjoint, nonempty subsets that together cover the allowed keys currently being considered. Put their maxima into the common coordinate system \(m\), define two scale factors, and combine the denominator and vector separately.
Expanding the old weighted sum shows why this works: the two occurrences of \(m_A\) cancel inside the exponents. The denominator follows the same transformation.
The new tile has the corresponding expression with the same \(m\). Thus \(u\) and \(\ell\) are the exponential weighted sum and denominator over the union. Dividing once yields the global output. Repeating this invariant over any number of tiles accounts for all allowed keys, rather than selecting a tile-local softmax.
In exact real arithmetic, merging disjoint subsets is associative and permits exchanging their order, because every merge tree represents the same set. This does not excuse repeated keys. Including a key twice counts its contribution twice in both the denominator and weighted sum. Summation order still affects finite-precision rounding. “Exact attention” means no sparse or low-rank approximation to the operator; it does not promise bitwise equality across backends.
A subtler bug rescales the denominator but forgets \(u\). Keep logits \((0,\log 3)\) and change the values to \((1,0)\). The correct output is \(1/4\). The new maximum gives \(\alpha=1/3\), which must multiply both the old denominator and weighted sum. Rescaling only the denominator produces \(3/4\). A plausible denominator is insufficient evidence of a correct answer.
5. From one row to tensor tiles: divide once at the end
An implementation usually processes \(B_q\) queries and \(B_k\) keys at a time. The score tile has shape \(B_q\times B_k\); every query keeps its own \(m,\ell,u\). Matrix products form scores and value-weighted sums, while row reductions update statistics. Rescaling broadcasts along rows: a single maximum shared by the entire tile would be wrong.
Writing the current tile's exponentials directly relative to the new maximum gives this convenient sequential update. Here \(B\) includes only the allowed keys in that tile for the current row.
On a nonempty row's first update, the old state has zero mass; there is no need to evaluate an empty state's exponentials literally. Later updates preserve the invariant. Only after the loop do we compute \(o=u/\ell\), avoiding repeated division of an entire output vector and reversal of that normalization at the next tile. Section 3.1 of FlashAttention-2 retains an unnormalized accumulator until the end; the paper also rearranges query-block parallelism and work among warps to reduce non-matmul operations and shared-memory communication. A valid recurrence and efficient GPU scheduling are separate things to establish.
for each query tile:
initialize per-row m, l, u as empty states
for each key/value tile:
compute scaled scores and the global-position mask
for each row with valid keys in this tile:
choose the new maximum
rescale BOTH the old l and old u
add this tile's mass and weighted values
normalize each nonempty row once; handle empty rows explicitly
This is an algorithm outline, not a CUDA kernel. The same loop expressed as several Python tensor operations can still write intermediates to device memory, repeatedly launch kernels, or retain every tile through automatic differentiation. Being mathematically online does not establish FlashAttention's memory traffic. Tile size cannot grow indefinitely either: registers, shared memory, and concurrent resident work blocks constrain usable parallelism.
6. Follow global positions, and handle empty tiles explicitly
In square causal self-attention, a key's global position cannot be later than the query's. Local row and column indices in separate query and key tiles cannot simply be compared: the tiles may come from different sequence positions. Entirely masked tiles can be skipped for efficiency. Treating every tile as if it started at position zero changes the visible history.
The second, rectangular bottom-right condition is an interface convention rather than a universal framework rule. The matrix indices \(i,j\) are zero-based, and the two sequences' ending positions align. When the query sequence is no longer than the key sequence, queries correspond to its last \(N_q\) positions. The README at the pinned official repository commit records that FlashAttention 2.1 changed unequal-length causal masking to bottom-right alignment. For \(N_q=2,N_k=5\), the two rows allow the first four and five keys. Top-left alignment allows only the first one and two, naturally changing the outputs.
When \(N_q>N_k\), bottom-right alignment leaves early rows empty. Padding can also mask a whole row. The softmax denominator above is undefined on an empty set. Filling the row with negative infinities and mechanically subtracting negative infinity produces NaN. Our code explicitly returns zero for an entirely empty row. The documented interface above also specifies zero output, but that cannot be generalized to every framework, loss, or backend.
If a tile has no allowed keys for a row, skip that row's update and preserve any existing valid state. Treating an empty state as a merge identity requires an explicit branch rather than evaluating \(\exp(-\infty-(-\infty))\). We disable dropout. With training dropout, backward recomputation must reproduce the forward random mask and scaling. Rotary positions, sliding windows, biases, and padding need the same alignment audit; a roughly similar mask applied after merging is insufficient.
7. Why training need not save the full probability matrix
If forward does not retain \(P\), can training still obtain gradients? Save each nonempty row's log-sum-exp \(L_i\) and output. Given the original Q/K and fixed mask, backward can reconstruct scores and then probabilities tile by tile. For loss \(\mathcal L\), let the upstream gradient at row \(i\) be \(g_i\), with rows forming \(G\), and let \(D\) be the gradient with respect to the scaled logits.
The scalar \(c_i\) comes from the softmax Jacobian. It initially appears as \(\sum_j p_{ij}(g_i^\top v_j)\); exchanging the sum and dot product gives \(g_i^\top o_i\). That supplies the Jacobian's row reduction without retaining the whole probability row. After reconstructing a tile of \(P,D\), accumulate Q/K/V gradients and discard the tile. Q/K gradients must still include the original scale.
This derivation covers nonempty rows with dropout disabled and fixed masks and biases. A trainable bias needs its own gradient, and RoPE or projection layers producing Q/K need their subsequent chain rules. We did not execute autograd or low-precision backward checks, and do not interpret the custom empty-row zero output as the derivative of standard softmax on an empty set. Appendix B of FA1 provides the full training algorithm including scaling, masks, and dropout.
Recomputation adds arithmetic but avoids storing and moving large matrices. Increased operation count and lower wall-clock time can therefore coexist; whether they do depends on device and workload. An implementation that saves every probability tile in backward, or changes the forward mask or random state during recomputation, has not established the path described here.
8. Which bytes disappear, and which pairs remain?
For a single dense head, leading arithmetic remains \(O(N_qN_k(d_k+d_v))\), or \(O(N^2d)\) for square same-dimension attention. A causal mask roughly halves allowed pairs without changing the order. Online softmax does not make every key's influence on each query a linear-time computation. It mainly avoids fully materializing \(N_qN_k\) score/probability intermediates in large device memory.
A straightforward tile workspace has the following element count for its score tile, queries, keys/values, and output accumulator. Constants additionally depend on double buffering, registers, dtypes, and the kernel, so this expression is not a claim about whole-device peak memory.
Complete input/output storage and a total of \(O(N_q)\) scalar statistics across all query rows must also be accounted for. Training adds gradients, optimizer state, and activations in other layers. “Linear additional state” must identify that it excludes inputs and outputs; the whole model does not fit in a handful of scalars.
Theorem 2 of FA1 compares HBM traffic in an idealized memory hierarchy for square same-dimension attention. Here \(M\) is fast-memory capacity in numerical elements, not unconverted bytes. In the stated range:
With fixed \(M,d\), this IO expression still grows quadratically in \(N\). It describes a reduction in traffic relative to available fast memory. Other scaling regimes require \(M\) to grow with \(N\). The theorem is not a measured bandwidth model for every modern kernel or GPU, nor a guarantee of equal speed benefits at every sequence length.
Prefill usually has \(N_q\approx N_k=N\), allowing tile reuse of Q/K/V while avoiding large intermediates. Single-token decode has \(N_q=1,N_k=L\): there was only one probability row to begin with, and reading a long KV cache is often central. This step involves \(O(L(d_k+d_v))\) pair arithmetic and retains a cache of \(O(L(d_k+d_v))\). Section 2.2 of the official README describes splitting KV loading for short queries. Mergeable states justify combining the split computations mathematically; extra launches and synchronization still require measurement.
This also clarifies the relationship to MLA cache compression. Online reduction changes the execution of attention intermediates, without automatically compressing historical KV representations. Combining mechanisms must preserve both sets of tensor and scale conventions. The original FA4 report of March 5, 2026 further examines exponential computation, shared memory, and pipelines under asymmetric hardware scaling. Removing one bottleneck does not eliminate the next.
9. Testable boundaries: what was actually checked here?
Download the complete CPU check script. It uses only the Python standard library and a fixed random seed. The same Q/K/V, scale, and explicit Boolean mask, with additional biases disabled, go through dense and online forward paths. Dense attention retains scores as a reference; online attention computes and merges blocks. Storage of test inputs, reference matrices, and reports is not a production GPU kernel's peak-memory measurement.
The executed environment was CPython 3.12.14 / binary64 (53-bit significand). Element-wise acceptance uses \(|x-y|\le 2\times10^{-12}+2\times10^{-12}\max(|x|,|y|)\). This tolerance applies to these small checks, not all devices and dtypes. All 28 cases and 59 comparisons passed. The shape column lists query count, key count, key dimension, and value dimension in that order. Error is the maximum absolute difference against the same-input dense reference within the family. No runtime or GPU-memory numbers are reported.
Check family
Q/K/V dimensions
Maximum absolute error
Noncausal, irregular tiles
5 / 9 / 4 / 3
2.22045e-16
Square global causal mask
7 / 7 / 5 / 2
2.22045e-16
Rectangular bottom-right, empty prefix rows
2 / 6 / 1 / 1 5 / 2 / 1 / 1
0
Padding, empty tiles and rows
2 / 4 / 3 / 2 4 / 7 / 3 / 4
2.22045e-16
Empty key axis
2 / 0 / 3 / 2
0
Common logit shifts and wide span
1 / 5 / 1 / 1 1 / 5 / 1 / 2
4.44089e-16
All six three-tile orders
3 / 6 / 4 / 3
2.22045e-16
16 fixed-seed shapes
16 shapes in script
4.44089e-16
The checks validate complete, nonduplicated key partitions, allowing empty and irregular tiles and reordered visitation. Padding, entirely empty rows, and intermediate tiles with no allowed keys test zero output or preservation of existing state. Extreme finite-logit cases test normalization after maximum subtraction, not arbitrary overflowing QK dot products. Two wrong implementations also produced the counterexamples in Sections 2 and 4. These are educational forward-algorithm checks, rather than FlashAttention kernel tests, gradient checks, training-quality evaluations, or serving speedups.
A deployment comparison should have three acceptance layers. First freeze weights and inputs to inspect outputs, masks/positions, dtypes, and gradients, covering unequal lengths, empty rows, padding, dropout, and extreme logits. Next hold GPU, backend version, batch/head dimensions, lengths, dtype, and causal settings fixed; warm up and synchronize before measuring attention kernels, prefill, decode, peak memory, and failures separately. Only then compare end-to-end latency, throughput, and training/task outcomes under explicit budgets and baselines. An API already selecting a fused backend should not be labeled an unoptimized materialized baseline. These GPU and model-level experiments are proposed, not executed here.
If changing tile size significantly changes the output, investigate scaling, row broadcasting, masks, and state rescaling first. If numerical checks pass but speed does not improve, inspect whether the backend truly fuses operations, whether tiles return to HBM, parallelism, cache layout, and synchronization. The former invalidates implementation correctness; the latter invalidates a speed claim for that workload. Distinguishing them tells us whether to repair the mathematics or the system.
The transferable method is to identify a mergeable state, then make every step's input, output, invariant, and failure conditions explicit. FlashAttention preserves the mathematical definition of global attention while changing where intermediates live and move. Avoiding a large matrix does not remove all pair interactions, and does not establish cache, gradient, or delivery costs by itself.
Citation metadata were checked with citation-management tooling; its software attribution is Scientific Agent Skills. This tooling citation is not evidence for the attention mechanism or performance conclusions.