FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
Dao et al., 2022 — computes exact attention in small tiles that stay in fast on-chip memory, never writing the full score matrix out, making long context faster and far cheaper in memory.
"FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness" (Tri Dao, Daniel Fu, Stefano Ermon, Atri Rudra, and Christopher Ré at Stanford and the University at Buffalo, 2022) introduced the kernel that most modern transformer training and serving now uses for attention.
What problem it solved
Self-attention builds an n × n score matrix, so cost grows with the
square of the sequence length, as the napkin math in
Attention & Transformers shows. Many
proposals tried to cut that cost with approximations: sparse patterns,
low-rank sketches, and so on. Those reduced the number of arithmetic
operations, but they often didn't make real models run faster, and they
changed the answer.
The paper's diagnosis was that arithmetic wasn't the bottleneck. On a modern GPU, moving data between the big, slow main memory (HBM) and the small, fast on-chip memory (SRAM) costs far more time than the multiplications do. Standard attention writes the whole score matrix out to HBM and reads it back several times, and that traffic, not the math, is where the time goes.
The key idea
Make attention IO-aware: count memory movement, not just operations,
and design the algorithm to minimize it. FlashAttention splits the queries,
keys, and values into blocks small enough to fit in SRAM and processes
them tile by tile, fusing the whole attention computation into one
kernel. The obstacle is the softmax, which normally needs an entire row of
scores before it can normalize. The paper uses an online softmax: keep
a running maximum and running sum for each row and rescale the partial
results as each new tile arrives, so the final answer is exact. The full
n × n matrix is never stored; in the backward pass it's recomputed from
the tiles rather than saved, trading a little extra arithmetic for a lot
less memory traffic.
Why it mattered
It was exact, not approximate, so it could replace standard attention
without changing a model's outputs, and it was faster in wall-clock time,
which most approximations weren't. The paper reported roughly 15% faster
end-to-end training on BERT-large and about 3x on GPT-2, with memory use
growing linearly instead of quadratically in sequence length. That
linear memory is one of the main reasons context windows jumped from a couple
of thousand tokens to hundreds of thousands: the n² matrix that made
long context look impractical simply stopped being stored. It also
reframed how people optimize deep learning code, putting memory
bandwidth, the same bottleneck behind slow decode in
Inference & Serving, at the center. The
mechanism is still quadratic in compute; what changed is that compute
wasn't the part hurting.
Authors: Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré (Stanford University, University at Buffalo)
Read the paper — arXiv:2205.14135Learn more: AttentionAttention (Self-Attention)Attention is a mechanism letting each position in a sequence weigh every other position via learned Query/Key/Value vectors, forming the core of the transformer. · Attention & Transformers · Inference & Serving
Retrieval-Augmented Generation for Knowledge-Intensive NLP Tasks
Lewis et al., 2020 — paired a text generator with a searchable index of documents so answers come from retrieved passages rather than only from what the model memorized.
Fast Inference from Transformers via Speculative Decoding
Leviathan, Kalman & Matias, 2022 — lets a small draft model guess several tokens ahead and a large model check them all at once, speeding up generation without changing a single output.