FlashAttention

Turkish equivalent: FlashAttentionDomain: Machine Learning Systems

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.

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.