Rounding error in FlashAttention: centered, uncentered, and radial

We bound the forward rounding error of FlashAttention-2 for one query and compare it with two-pass attention. The analysis rests on one identity. An error that hits a key's weight in both the numerator and the normalizer costs only the spread of the values around the output. Most of FlashAttention's roundings are of this kind. What is left is the cast of the weights to the input format, and both algorithms pay the same worst-case price for it. On peaked attention FlashAttention-2 does better, because in PyTorch's build the key at the running maximum has weight exactly 1 and casts without error. The split-KV merge breaks the identity. Its weights are never renormalized, so the rounding of the global log-sum-exp scales the whole output, and the error grows with the logits. On an NVIDIA L4, constant values that should give exactly 1 come back up to 2.2% off. Renormalizing the merge removes this error.

Authors

Publication Details

Journal
Zenodo (CERN European Organization for Nuclear Research)
Published
2026-10-06
DOI
https://doi.org/10.5281/zenodo.23179059
Primary Topic
Numerical Methods and Algorithms
Type
preprint
Controls
|||
ALL TIME
JAN
FEB
MAR
APR
MAY
JUN
JUL
AUG
SEP
OCT
preprint

Rounding error in FlashAttention: centered, uncentered, and radial

Hanyu Yang
Zenodo (CERN European Organization for Nuclear Research)
Numerical Methods and Algorithms
preprint

Rounding error in FlashAttention: centered, uncentered, and radial

Hanyu Yang
preprint en

Abstract

We bound the forward rounding error of FlashAttention-2 for one query and compare it with two-pass attention. The analysis rests on one identity. An error that hits a key's weight in both the numerator and the normalizer costs only the spread of the values around the output. Most of FlashAttention's roundings are of this kind. What is left is the cast of the weights to the input format, and both algorithms pay the same worst-case price for it. On peaked attention FlashAttention-2 does better, because in PyTorch's build the key at the running maximum has weight exactly 1 and casts without error. The split-KV merge breaks the identity. Its weights are never renormalized, so the rounding of the global log-sum-exp scales the whole output, and the error grows with the logits. On an NVIDIA L4, constant values that should give exactly 1 come back up to 2.2% off. Renormalizing the merge removes this error.

Zenodo (CERN European Organization for Nuclear Research)
Numerical Methods and Algorithms
AI Navigator

Ask Laika to Summarize, Analyze, and Connect papers live on the map.

Summarize Papers & Methodologies

Extract key findings, datasets, and comparative methods across publications.

Benchmark Rankings & Visual Analytics

Rank top research institutions, authors, funders, topics, and journals by Field-Weighted Citation Impact (FWCI) and paper volume with instant charts.

Connect Distant Disciplines

Bridge topological clusters on the map to find hidden collaborative intersections.