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 softmax(QK⊤/dk) V\text{softmax}(QK^\top / \sqrt{d_k})\,V. For a sequence of NN tokens this creates an N×NN \times N score matrix — that's N2N^2 numbers to store and process. At N=4096N = 4096, 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 N×NN \times N 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.

Open in Lab
Click each memory level to see its size and speed. Notice the 100× gap between SRAM and HBM bandwidth.
The demo wakes as you arrive…

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 N×NN \times N 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 N×NN \times N 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.

Open in Lab
Watch how FlashAttention tiles through Q and K/V blocks. The full N×N matrix is never stored — only small block results accumulate in SRAM.
The demo wakes as you arrive…

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 softmax(xi)=exi/∑jexj\text{softmax}(x_i) = e^{x_i}/\sum_j e^{x_j} requires knowing all scores xjx_j to get the denominator. How can we tile something that needs to see the whole row?

The insight: maintain running statistics — a running maximum mm and a running sum-of-exponentials ℓ\ell. When a new block of scores arrives, update mm and ℓ\ell, 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.

m(new)=max⁡(m(old), m~),ℓ(new)=em(old)−m(new)ℓ(old)+em~−m(new)ℓ~m^{(new)} = \max(m^{(old)},\, \tilde{m}), \qquad \ell^{(new)} = e^{m^{(old)} - m^{(new)}} \ell^{(old)} + e^{\tilde{m} - m^{(new)}} \tilde{\ell}
Online softmax rescaling — updating running statistics when a new block arrives — m = running row-maximum · ℓ = running sum of exponentials · tilde = stats from the new block · the exponential factors correct for the old maximum being stale
Open in Lab
Step through the online softmax: see how the running max and sum update with each new block, and verify the final result matches full-row softmax.
The demo wakes as you arrive…

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 S=QK⊤S = QK^\top, the softmax result PP, and often the dropout mask. These are all N×NN \times N matrices. FlashAttention keeps all intermediates in SRAM and only writes the final output OO (which is N×dN \times d, not N×NN \times N) to HBM.

Memory usage drops from O(N2)O(N^2) to O(N)O(N). 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.

Open in Lab
Drag the sequence length slider and compare memory usage: standard attention (quadratic) vs. FlashAttention (linear).
The demo wakes as you arrive…

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 Θ(Nd+N2)\Theta(Nd + N^2) HBM accesses. FlashAttention requires Θ(N2d2M−1)\Theta(N^2 d^2 M^{-1}) where MM is SRAM size. Since MM is large relative to d2d^2 (typical d=64–128d = 64\text{–}128, typical M=100KB+M = 100\text{KB+}), 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.

HBM accessesstandard=Θ(Nd+N2)HBM accessesflash=Θ ⁣(N2d2M)\text{HBM accesses}_{\text{standard}} = \Theta(Nd + N^2) \qquad \text{HBM accesses}_{\text{flash}} = \Theta\!\left(\frac{N^2 d^2}{M}\right)
IO complexity comparison — standard vs. FlashAttention — N = sequence length · d = head dimension · M = SRAM size · Flash wins because M ≫ d², making the denominator large

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.

  1. Divide Q into blocks of size BrB_r, and K, V into blocks of size BcB_c.
  2. Initialize output O=0O = 0, running max m=−∞m = -\infty, running sum ℓ=0\ell = 0.
  3. For each K/V block jj: load Kj,VjK_j, V_j into SRAM.
  4. — For each Q block ii: load QiQ_i into SRAM, compute Sij=QiKj⊤S_{ij} = Q_i K_j^\top.
  5. — — Compute local max m~\tilde{m} and local sum ℓ~\tilde{\ell} from SijS_{ij}.
  6. — — Update running mm and ℓ\ell using the online softmax formulas.
  7. — — Rescale the running output OiO_i and add the new block's contribution.
  8. Write the final OO and save ℓ,m\ell, m (for the backward pass) to HBM.
Open in Lab
Step through the FlashAttention algorithm block by block. Watch the running statistics update and the output accumulate — no N×N matrix ever forms.
The demo wakes as you arrive…

The same idea in code

FlashAttention forward pass (simplified Python)python

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 (i,j)(i, j) 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).

Open in Lab
Toggle sparsity patterns and see which blocks FlashAttention skips. Sparser patterns mean fewer blocks to compute and fewer HBM trips.
The demo wakes as you arrive…

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

  1. 2018

    Online softmax

    Milakov & Gimelshein introduce online softmax for stable, single-pass computation — the mathematical foundation FlashAttention builds on.

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

  3. 2023

    FlashAttention-2

    Better parallelism and work partitioning bring FlashAttention closer to hardware limits, reaching 50–73% of theoretical FLOPS on A100.

  4. 2024

    FlashAttention-3

    Leverages Hopper GPU features (asynchronous memory, FP8) to push throughput even higher, approaching hardware peak on H100.

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