Rounding error in FlashAttention: centered, uncentered, and radial

A forward error bound for FlashAttention-2, a comparison with two-pass attention, and a measured scale error in the split-KV merge.

DOI 10.5281/zenodo.23179058 · PDF · alphaXiv · illustrated summary · CC BY 4.0 · preprint, 6 Oct 2026 · ORCID 0009-0005-0419-4070
FlashAttention-2 and two-pass attention share the same leading worst-case term, the cast of the weights to bf16 or fp16. FlashAttention-2 is more accurate on peaked attention. The split-KV merge never renormalizes its weights, so rounding the global log-sum-exp scales the whole output: on an NVIDIA L4, an output that should be exactly 1 comes back up to 2.2% off. Renormalizing the merge removes the error.

The idea

Perturb each softmax weight by a factor 1+βj. The output moves by Σj pj βj (vj − y), up to a normalizing factor. So an error that hits a key's weight in the numerator and the normalizer alike costs only the spread of the values around the output. That sorts every rounding in the kernel into three kinds.

Three panels: centered errors cost the spread of the values, uncentered errors cost their size, radial errors scale the output.
Centered, uncentered and radial errors, and where each one comes from.

Measured on an NVIDIA L4

PyTorch 2.11 FlashAttention-2, fp16, d = 64, τ = 1/8, identical keys and every value equal to 1, so the exact output is 1. Largest |ŷ − 1| in each logit range, in units of 2−11:

batch × keyskernelτs < 27[27, 212)[212, 216)τs ≥ 216
24 × 256single pass0000
24 × 1024split-KV, T = 400≤ 1≤ 24
24 × 2048split-KV, T = 400≤ 1≤ 24
1 × 4096split-KV, T = 1600≤ 2≤ 46
Measured error against logit scale for four shapes: single pass stays at zero, split-KV rises in steps to 24 and 46 units of 2 to the minus 11.
All 160 logit levels per shape. The shaded region is where the split-KV bound is vacuous. The raw 640 measurements are on Zenodo as CSV.

A published bound that fails

The O(u log n) relative-error bound of Hsu et al. (ELSA, Theorem 1) for online-softmax attention is false as stated. It misses a cancellation factor and a logit-size term; the paper gives counterexamples of both kinds, run in their pairwise merge order against an 80-digit reference.

What it does not establish

The fix itself is not new: FlashAttention-3, FlashAttention-4, FlashInfer, vLLM and xformers already renormalize their merges. The paper is AI-assisted; its disclosure section says how.