FlashAttention
An exact attention algorithm that reorganizes computation around fast on-chip memory to reduce memory traffic and improve attention throughput.
ML-Systems Context
FlashAttention computes exact softmax attention using a tiled algorithm that keeps more intermediate work in fast on-chip memory and reduces transfers to high-bandwidth memory. Long-sequence training and inference can therefore use less memory traffic and achieve higher throughput.
Algorithm Boundary
FlashAttention is not an approximate-attention method; it reorganizes the computation of the same attention result. The practical gain depends on GPU architecture, sequence shape, datatype, and kernel implementation.
Related ML-Systems Concepts
Data Movement Before Arithmetic Count
FlashAttention is distinctive because it treats data movement through the GPU memory hierarchy as a first-class cost, rather than considering FLOPs alone. Attention is tiled so that more intermediate work stays in fast on-chip SRAM and fewer reads/writes go to HBM. The same mathematical softmax attention is therefore evaluated with a different execution schedule.
The word "exact" matters: the core FlashAttention algorithm is not a low-rank or sparse approximation to attention. Exactness at the mathematical algorithm level does not imply bit-for-bit identical floating-point output, because reduction order and datatype can still change rounding. Practical speedup depends on sequence length, head dimension, datatype, GPU architecture, and kernel version. Dao et al. define the I/O-aware exact-attention approach directly: original paper.