Systems & Optimization2022advanced10 min read
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
FlashAttention: انتباه دقيق وسريع وموفِّر للذاكرة بوعي حركة البيانات
Dao, T. · Fu, D. Y. · Ermon, S. · Rudra, A. · Ré, C. — NeurIPS
The problem
Self-'s time and memory scale quadratically with sequence length: the N×N attention matrix must be materialized in GPU high-bandwidth memory (). Approximate methods trade accuracy for speed but often fail to deliver real wall-clock speedups because they ignore the true bottleneck — data movement between slow HBM and fast on-chip , not raw .
The contribution
An IO-aware algorithm that uses to load small blocks of Q, K, V into fast SRAM, computes attention block-by-block using , and never materializes the N×N attention matrix in HBM. merges five separate operations into one GPU kernel. replaces storage: the recomputes attention on the fly instead of caching the full matrix. Result: 2–4× , 10–20× memory savings, and the same exact mathematical result as standard attention.
The impact
FlashAttention became the de facto standard for attention in production Transformers. GPT-4, Claude, LLaMA, and virtually every modern large language model uses FlashAttention or its successors (FlashAttention-2, FlashAttention-3). It enabled context windows to grow from 2K to 128K+ tokens. Its IO-aware philosophy shifted the community's focus from FLOP reduction to memory-bandwidth optimization — a paradigm change in systems-level ML research.
Imagine a chef preparing a banquet. The standard recipe says: carry all the ingredients from the warehouse to the kitchen counter, lay them out, cook, then carry the results back. The counter is tiny, so most ingredients wait in the warehouse — and the chef spends more time walking back and forth than actually cooking.
FlashAttention's recipe: bring a small tray of ingredients at a time, cook them immediately on the counter, jot down a running note of what's done, and repeat. The full spread never needs to exist all at once. The chef barely leaves the kitchen, and the banquet is served in half the time — with exactly the same dishes.
The real bottleneck: memory bandwidth, not compute
The Transformer computes attention as . For a sequence of tokens this creates an score matrix — that's numbers to store and process. At , that's 16 million entries per head, per layer.
But here's the surprise: for typical sequence lengths, the bottleneck is not the arithmetic. Modern GPUs can do trillions of floating-point operations per second (FLOPS). What they cannot do is move data fast enough. The A100 GPU computes at 312 TFLOPS but its HBM bandwidth is only 2 TB/s. If the ratio of compute to data movement is low — as it is for element-wise operations like and masking — the GPU cores sit idle, waiting for data. This is what systems researchers call a workload.
Standard attention writes the full matrix to HBM, reads it back for softmax, writes the softmax result, and reads it again for the final matrix multiply. Each of those round-trips to slow memory is wasted time. Approximate methods (sparse attention, linear attention) tried to fix this by reducing FLOPs — but since FLOPs were never the bottleneck, they often failed to speed things up in practice.
Three ideas, one kernel
FlashAttention combines three classical systems techniques — tiling, kernel fusion, and recomputation — into a single GPU kernel that computes exact attention without ever writing the matrix to HBM.
1. Tiling — divide Q, K, V into small blocks that fit in SRAM. Load one block of K and V at a time, compute partial attention for all Q blocks against that K/V block, then move to the next. The full attention matrix is never assembled.
2. Kernel fusion — standard attention runs five separate GPU kernels: matmul → mask → softmax → → matmul. Each kernel reads from and writes to HBM. FlashAttention fuses all five into one kernel: data enters SRAM once, all five operations happen there, and only the final output leaves.
3. Recomputation — in the backward pass, instead of storing the attention matrix for computation, FlashAttention recomputes it on the fly from Q, K, V (which are already stored). This trades a small amount of extra FLOPs for massive memory savings.
The online softmax trick: the key insight
Tiling works naturally for matrix multiplication because it's associative — you can split it into blocks and combine results. But softmax is not directly associative: computing requires knowing all scores to get the denominator. How can we tile something that needs to see the whole row?
The insight: maintain running statistics — a running maximum and a running sum-of-exponentials . When a new block of scores arrives, update and , then rescale the partially accumulated output by the correction factor. After processing all blocks, the result is mathematically identical to computing softmax over the full row at once.
This is called online softmax (Milakov & Gimelshein, 2018). It allows the softmax computation to be split into arbitrarily many blocks, each processed independently in SRAM, with only a small running state carried between blocks.
Standard attention vs. FlashAttention: the memory story
The fundamental difference is where computation happens. Standard attention materializes every intermediate result in HBM — the score matrix , the softmax result , and often the dropout mask. These are all matrices. FlashAttention keeps all intermediates in SRAM and only writes the final output (which is , not ) to HBM.
Memory usage drops from to . At sequence length 2K, that's 10× savings. At 4K, it's 20×. This is what lets modern models use 128K+ context windows — something physically impossible with standard attention on available GPUs.
IO complexity: why fewer trips matter more than fewer FLOPs
FlashAttention's key theoretical contribution is analyzing attention through the lens of — counting HBM reads and writes instead of floating-point operations.
Standard attention requires HBM accesses. FlashAttention requires where is SRAM size. Since is large relative to (typical , typical ), the number of HBM accesses is much smaller.
The authors also prove a lower bound: for any exact attention algorithm using standard operations, FlashAttention's IO complexity is optimal up to constant factors for a range of SRAM sizes. You literally cannot do better with this hardware model.
The algorithm step by step
The works as follows. Before seeing the pseudocode, understand the rhythm: outer loop over K/V blocks, inner loop over Q blocks — for each pair, compute a small tile of attention, update running softmax statistics, and accumulate into the output.
- Divide Q into blocks of size , and K, V into blocks of size .
- Initialize output , running max , running sum .
- For each K/V block : load into SRAM.
- — For each Q block : load into SRAM, compute .
- — — Compute local max and local sum from .
- — — Update running and using the online softmax formulas.
- — — Rescale the running output and add the new block's contribution.
- Write the final and save (for the backward pass) to HBM.
The same idea in code
Simplified to show the idea — not the real implementation.
import numpy as np
def flash_attention(Q, K, V, block_size=64):
"""FlashAttention: tiled exact attention without materializing N×N."""
N, d = Q.shape
O = np.zeros_like(Q) # output accumulator
ell = np.zeros((N, 1)) # running sum of exponentials
m = np.full((N, 1), -np.inf) # running row-max
# Outer loop: stream K/V blocks (like loading trays into the kitchen)
for j in range(0, N, block_size):
Kj = K[j:j+block_size]
Vj = V[j:j+block_size]
# Inner loop: each Q block processes this K/V block
for i in range(0, N, block_size):
Qi = Q[i:i+block_size]
Sij = Qi @ Kj.T / np.sqrt(d) # small tile of scores
# Online softmax: update running max and sum
m_new = np.maximum(m[i:i+block_size], Sij.max(axis=-1, keepdims=True))
P = np.exp(Sij - m_new) # safe exp with new max
correction = np.exp(m[i:i+block_size] - m_new)
# Rescale old output and add new contribution
O[i:i+block_size] = correction * O[i:i+block_size] + P @ Vj
ell[i:i+block_size] = correction * ell[i:i+block_size] + P.sum(axis=-1, keepdims=True)
m[i:i+block_size] = m_new
return O / ell # final normalization
# That's it. The N×N matrix never exists.
# Real CUDA kernels do this in on-chip SRAM, not main memory.Block-sparse FlashAttention: skipping what doesn't matter
FlashAttention naturally extends to : if a sparsity mask says that block is zero, simply skip it entirely — no load, no compute, no write. This is trivial in the tiled framework because each block is already a distinct unit of work.
Block-sparse FlashAttention yields an additional 2–4× speedup on top of dense FlashAttention, proportional to the sparsity ratio. This enabled with sequences of up to 64K tokens — and the first Transformer to achieve better-than-chance on the Path-256 benchmark (sequence length 64K, 63.1% accuracy).
Results: faster training, longer contexts, better models
FlashAttention's impact is both in speed and in what longer contexts enable:
- 15% end-to-end speedup on BERT-large (sequence length 512) compared to the MLPerf 1.1 training speed record.
- 3× speedup on GPT-2 (sequence length 1K) in attention computation.
- 2.4× speedup on long-range arena tasks (sequence lengths 1K–4K).
- 0.7 better perplexity on GPT-2 with longer context (better modeling, same model size).
- 6.4 point lift on long-document classification tasks.
- First better-than-chance on Path-X (16K tokens, 61.4%) and Path-256 (64K tokens, 63.1%) — benchmarks no Transformer could previously handle.
The legacy and what came after
2018
Online softmax
Milakov & Gimelshein introduce online softmax for stable, single-pass computation — the mathematical foundation FlashAttention builds on.
2022
FlashAttention (this paper)
Dao et al. combine tiling, kernel fusion, and recomputation to achieve exact attention with 2–4× speedup and linear memory. Published at NeurIPS 2022.
2023
FlashAttention-2
Better parallelism and work partitioning bring FlashAttention closer to hardware limits, reaching 50–73% of theoretical FLOPS on A100.
2024
FlashAttention-3
Leverages Hopper GPU features (asynchronous memory, FP8) to push throughput even higher, approaching hardware peak on H100.
2024
128K+ context windows
Claude, GPT-4 Turbo, Gemini 1.5, and other models ship with 128K–1M token context windows, all made feasible by FlashAttention-family algorithms.
FlashAttention's lesson extends beyond attention. Any operation that is memory-bound — , dropout, activation functions — benefits from the same IO-aware philosophy. The paper changed how the ML systems community thinks about optimization: start with the , not the FLOP count.
CitationDao, Fu, Ermon, Rudra, Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS, 2022.
Terms in this paper
- High Bandwidth Memory (HBM)الذاكرة عالية النطاق
- SRAMالذاكرة الساكنة
- Tilingالتبليط
- Kernel Fusionدمج النواة
- Recomputationإعادة الحساب
- Online Softmaxsoftmax المتدفق
- IO-Awarenessالوعي بحركة البيانات
- Memory-Boundمقيَّد بالذاكرة
- Compute-Boundمقيَّد بالحساب
- Block-Sparse Attentionالانتباه المتناثر الكُتَلي
- Wall-Clock Speedupتسريع الزمن الفعلي
- IO Complexityتعقيد المدخلات والمخرجات
- Materializationالتجسيد
- Exact Attentionالانتباه الدقيق
- Memory Hierarchyهرمية الذاكرة