Inference & Efficiency

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.14135

Learn more: Attention · Attention & Transformers · Inference & Serving

On this page