Fast Attention and Long Context

FlashAttention

FlashAttention gives the same answer as ordinary attention while never writing the giant score table to memory, which is what makes long contexts affordable.

Read these first

On this page 8
  1. The short answer
  2. The kitchen you have cooked in
  3. What was broken before
  4. The trick, in one idea
  5. Why the word "exact" matters
  6. Where you have already seen it
  7. Remember this
  8. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

FlashAttention is the attention you already know, rearranged so the computer never writes down the enormous middle step.

The kitchen you have cooked in

Picture a small kitchen. The fridge is across the room and the counter is right in front of you. Fetch one onion, walk back, chop it, then walk to the fridge again. You will spend your evening walking.

A good cook carries a whole tray over once. Everything happens on the counter, and only the finished dish goes back.

A graphics card has the same two places. It has a large slow store, and a tiny fast workspace beside each processing unit. FlashAttention is the cook who stopped walking.

What was broken before

Attention compares every word with every other word. Attention here means the step where the model decides which earlier words matter for the word it is writing now. If you have not met it yet, read attention first.

The ordinary way writes that comparison out in full. For a thousand words that is a table with a million entries. For a hundred thousand words it is ten billion entries.

That table is written to the slow store and read back. Nobody wants the table. It exists for a moment and is thrown away.

So the model spends most of its time moving a throwaway table around. The chopping is fast. The walking is slow.

The trick, in one idea

You do not need the whole table to get the answer. You need a running total — a summary you keep updating as you go. It is like adding up a shopping bill in your head.

FlashAttention brings over one strip of the table at a time. It updates the running total on the fast counter. Then it throws the strip away and fetches the next one.

   the old way                    FlashAttention
   -----------                    --------------
   compute the WHOLE table        take one strip
   write it to the slow store     work on it on the fast counter
   read it back                   update a running total
   finish                         throw the strip away, take the next
                                  finish

   answer: identical              answer: identical
   memory: enormous               memory: tiny

There is one wrinkle that makes this harder than a shopping bill. Attention turns scores into percentages, and a percentage depends on every score, including ones you have not seen yet.

This part is confusing for almost everyone the first time. Read it twice, that is normal.

The fix is a correction step. A later strip may contain a bigger score than anything seen so far. The running total is then scaled down to match. The bookkeeping is exact, so the final answer is not an approximation.

Why the word "exact" matters

Many long-context methods drop comparisons to save time. They give an answer that is close, not equal.

FlashAttention drops nothing. The output is the same output, to the last decimal the hardware can hold. That is why it spread everywhere within a year, with no quality argument to have.

Where you have already seen it

  • Every chatbot that accepts a long document you paste in.
  • Fine-tuning a model on a free cloud notebook without running out of memory.
  • Any local model server running on your own machine.

Remember this

  • FlashAttention never stores the full comparison table. It works one strip at a time.
  • The saving is in memory traffic, meaning data carried between the slow store and the fast workspace.
  • The answer is exact, not approximate. Nothing is skipped.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install numpy        # for the first example
pip install torch        # for the second; a CUDA GPU is needed for the measurement

Written and run against NumPy 1.26.4 and PyTorch 2.5.1+cu121 on Python 3.10.

The algorithm, in plain NumPy

The whole idea is the loop below. It keeps three running quantities and never builds an N x N array.

flash_numpy.py
import numpy as np

rng = np.random.default_rng(0)
N, d, B = 512, 32, 128          # 512 tokens, head dim 32, tiles of 128 keys

Q = rng.normal(size=(N, d)).astype(np.float32)
K = rng.normal(size=(N, d)).astype(np.float32)
V = rng.normal(size=(N, d)).astype(np.float32)
scale = 1.0 / np.sqrt(d)

def naive_attention(Q, K, V):
    S = (Q @ K.T) * scale                      # the N x N score matrix, held in full
    S = S - S.max(axis=-1, keepdims=True)
    P = np.exp(S)
    return (P / P.sum(axis=-1, keepdims=True)) @ V, S.nbytes

def flash_attention(Q, K, V, block=B):
    n = Q.shape[0]
    O = np.zeros_like(Q)                       # running output
    m = np.full(n, -np.inf, dtype=np.float32)  # running row max
    l = np.zeros(n, dtype=np.float32)          # running sum of exp
    biggest = 0
    for j in range(0, n, block):               # one tile of keys at a time
        Kj, Vj = K[j:j + block], V[j:j + block]
        Sj = (Q @ Kj.T) * scale                # N x block, never N x N
        biggest = max(biggest, Sj.nbytes)
        m_new = np.maximum(m, Sj.max(axis=-1))
        correction = np.exp(m - m_new)         # rescale what we already accumulated
        Pj = np.exp(Sj - m_new[:, None])
        l = correction * l + Pj.sum(axis=-1)
        O = correction[:, None] * O + Pj @ Vj
        m = m_new
    return O / l[:, None], biggest

out_naive, bytes_naive = naive_attention(Q, K, V)
out_flash, bytes_flash = flash_attention(Q, K, V)

print("outputs match:", np.allclose(out_naive, out_flash, atol=1e-5))
print("largest absolute difference:", float(np.abs(out_naive - out_flash).max()))
print()
print(f"naive: biggest score matrix held = {bytes_naive:,} bytes  ({N} x {N})")
print(f"flash: biggest score tile held   = {bytes_flash:,} bytes  ({N} x {B})")
print(f"ratio = {bytes_naive / bytes_flash:.0f}x smaller")
print()
for n in (1024, 8192, 131072):
    full = n * n * 2 / 1e9                     # bf16 scores for one head
    print(f"seq {n:>7}: one head's full score matrix in bf16 = {full:10.2f} GB")
Output
outputs match: True
largest absolute difference: 2.384185791015625e-07

naive: biggest score matrix held = 1,048,576 bytes  (512 x 512)
flash: biggest score tile held   = 262,144 bytes  (512 x 128)
ratio = 4x smaller

seq    1024: one head's full score matrix in bf16 =       0.00 GB
seq    8192: one head's full score matrix in bf16 =       0.13 GB
seq  131072: one head's full score matrix in bf16 =      34.36 GB

Reading that output

The difference is 2.4e-07. That is float32 rounding, not an approximation. This is the claim that made FlashAttention easy to adopt: no quality trade-off to argue about.

The ratio is 4x here because the tile is a quarter of the sequence. The real saving grows with sequence length, because the tile stays a fixed size while N grows. Memory goes from growing with the square of the length to growing with the length.

The last three lines are the reason anyone cares. At 131,072 tokens, one attention head's score matrix is 34 GB. A model has dozens of heads and dozens of layers. Without tiling, long context is not expensive, it is impossible.

The three running quantities

m is the largest score seen so far in each row. Subtracting it before exp stops the exponential from overflowing. This is the standard safe-softmax trick applied to a stream.

l is the running sum of exponentials, which is the softmax denominator.

O is the running weighted sum of value vectors.

correction = np.exp(m - m_new) is the line that earns the whole algorithm. When a new tile raises the row maximum, everything accumulated under the old maximum was scaled by the wrong constant. Multiplying by this factor repairs l and O together, exactly.

Measuring it on a real GPU

PyTorch hides several attention kernels behind one function. MATH is the textbook version that builds the score matrix. The tiled kernels do not.

flash_torch.py
import torch, torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel

assert torch.cuda.is_available(), "this measurement needs a CUDA GPU"
print("gpu:", torch.cuda.get_device_name(0), "| torch:", torch.__version__)

B, H, N, D = 1, 16, 8192, 64
q, k, v = (torch.randn(B, H, N, D, device="cuda", dtype=torch.float16) for _ in range(3))

def works(backend):
    try:
        with sdpa_kernel(backend):
            F.scaled_dot_product_attention(q[:, :, :64], k[:, :, :64], v[:, :, :64])
        return True
    except RuntimeError:
        return False

for b in (SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION,
          SDPBackend.CUDNN_ATTENTION, SDPBackend.MATH):
    print(f"  backend {b.name:20s} usable here: {works(b)}")

def peak_mb(backend):
    torch.cuda.synchronize(); torch.cuda.empty_cache()
    torch.cuda.reset_peak_memory_stats()
    base = torch.cuda.memory_allocated()
    with sdpa_kernel(backend):
        out = F.scaled_dot_product_attention(q, k, v)
    torch.cuda.synchronize()
    mb = (torch.cuda.max_memory_allocated() - base) / 2**20
    return out, mb

o_math, mb_math = peak_mb(SDPBackend.MATH)
o_tiled, mb_tiled = peak_mb(SDPBackend.EFFICIENT_ATTENTION)

print()
print("same answer:", torch.allclose(o_math, o_tiled, atol=2e-3))
print(f"MATH backend      peak extra memory: {mb_math:8.1f} MB")
print(f"tiled backend     peak extra memory: {mb_tiled:8.1f} MB")
print(f"the N x N scores alone would be     {B*H*N*N*2/2**20:8.1f} MB in fp16")
Output
gpu: NVIDIA RTX A6000 | torch: 2.5.1+cu121
  backend FLASH_ATTENTION      usable here: False
  backend EFFICIENT_ATTENTION  usable here: True
  backend CUDNN_ATTENTION      usable here: True
  backend MATH                 usable here: True

same answer: True
MATH backend      peak extra memory:  13440.1 MB
tiled backend     peak extra memory:     16.0 MB
the N x N scores alone would be       2048.0 MB in fp16

PyTorch also prints several UserWarning lines to stderr explaining why each rejected kernel was rejected. Read them, they are unusually informative.

These numbers come from one machine: an NVIDIA RTX A6000 running Windows. Peak memory is deterministic for a fixed shape, so the same shapes reproduce the same figures. The list of usable backends will differ on your machine, which is exactly why the script prints it.

13,440 MB against 16 MB. The MATH path is worse than the raw 2,048 MB of scores because it promotes to float32 and holds several intermediates at once. The tiled path allocates the output and little else.

FLASH_ATTENTION reports False here. This Windows PyTorch wheel was not built with that kernel, and the card is Ampere rather than Hopper. The memory-efficient kernel uses the same tiling idea, so the lesson holds. Print the table rather than assuming: the answer depends on your operating system, GPU generation, dtype and head dimension.

Common mistakes

Expecting a speed-up on short sequences. At 128 or 512 tokens the score matrix fits in cache anyway. Tiled kernels win when the sequence is long, and can lose below a few hundred tokens.

Passing float32. The fastest kernels want float16 or bfloat16. Hand them float32 and PyTorch falls back to a slower path without complaining. Print the backend you actually got.

Passing a dense mask when is_causal=True would do. A boolean mask of shape N x N re-creates the exact allocation you were trying to avoid. Use the flag.

Installing flash-attn and assuming it is in use. pip install flash-attn --no-build-isolation needs a matching CUDA toolkit and a long compile. Check the backend per call, not the install log.

Reading a speed claim without its hardware. FlashAttention-2 targets Ampere through Ada. FlashAttention-3 is Hopper only. FlashAttention-4 targets Hopper and Blackwell. A benchmark from one generation says nothing about another.

Try it yourself

Change B in the NumPy script to 512, so there is one tile. Confirm the memory ratio becomes 1x and the answer does not move. Then set it to 16 and watch the correction step fire far more often. Print m inside the loop to see the row maximum climb.

What to learn next

Researcher — Mathematics and papers.

The problem being solved

Standard attention for one head, with $Q, K, V \in \mathbb{R}^{N \times d}$:

$$ O = \operatorname{softmax}!\left(\frac{QK^{\top}}{\sqrt{d}}\right) V $$

$N$ is sequence length, $d$ is head dimension, and $\sqrt{d}$ keeps the logits' variance stable at initialisation. The intermediates $S = QK^{\top}/\sqrt{d}$ and $P = \operatorname{softmax}(S)$ are each $N \times N$.

FLOPs are $\Theta(N^2 d)$ and cannot be reduced without changing the function. Memory traffic can be. Writing and re-reading $S$ and $P$ costs $\Theta(N^2)$ accesses to high-bandwidth memory (HBM), and on modern accelerators that traffic, not the arithmetic, sets wall-clock time.

Online softmax

The recurrence comes from Milakov and Gimelshein (2018), Online normalizer calculation for softmax (arXiv:1805.02867). Process key blocks $j = 1 \dots T$. After block $j$, hold $m^{(j)}$, the running row maximum, $\ell^{(j)}$, the running exponential sum, and $O^{(j)}$, the running unnormalised output.

$$ m^{(j)} = \max!\left(m^{(j-1)},\; \operatorname{rowmax}(S^{(j)})\right) $$

$$ \ell^{(j)} = e^{m^{(j-1)} - m^{(j)}} \ell^{(j-1)} + \operatorname{rowsum}!\left(e^{S^{(j)} - m^{(j)}}\right) $$

$$ O^{(j)} = e^{m^{(j-1)} - m^{(j)}} O^{(j-1)} + e^{S^{(j)} - m^{(j)}} V^{(j)} $$

$S^{(j)} \in \mathbb{R}^{N \times B_c}$ is the score block for key tile $j$ of width $B_c$, and $V^{(j)}$ is the matching value tile. The output is $O^{(T)} / \ell^{(T)}$. The factor $e^{m^{(j-1)} - m^{(j)}}$ rescales prior partial sums onto the new maximum, so the result matches the one-shot computation up to floating-point reassociation.

IO complexity

Dao, Fu, Ermon, Rudra and Ré (2022), FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (arXiv:2205.14135, NeurIPS 2022) prove HBM accesses of

$$ \Theta!\left(\frac{N^2 d^2}{M}\right) $$

against $\Theta(Nd + N^2)$ for the standard implementation, where $M$ is on-chip SRAM size in elements. For $d = 64$ and $M$ of order $10^5$ this is a 5–20x traffic reduction. They further show it is IO-optimal: no exact attention algorithm uses asymptotically fewer HBM accesses for all $M$.

Activation memory falls from $\Theta(N^2)$ to $\Theta(N)$, because $S$ and $P$ are never stored. The backward pass recomputes them tile by tile from the saved $\ell$ and $m$ statistics. That recomputation is close to free in wall-clock terms precisely because the kernel is memory-bound.

Rabe and Staats (2021), Self-attention Does Not Need $O(n^2)$ Memory (arXiv:2112.05682) reached the low-memory result independently. FlashAttention's contribution is the fused CUDA kernel that makes it faster rather than only smaller.

Version lineage

VersionYearTargetKey change
FlashAttention2022AmpereTiling plus recomputation; IO-optimal exact attention
FlashAttention-22023Ampere, Ada, HopperFewer non-matmul FLOPs, parallelism over sequence, better warp partitioning; about 2x over v1
FlashAttention-32024Hopper onlyWarp-specialised producer/consumer pipelining, FP8 forward
FlashAttention-42026Hopper, BlackwellCuTeDSL rewrite, fully asynchronous MMA, larger tiles, software-emulated exponentials
  • Dao (2023), FlashAttention-2 — arXiv:2307.08691, ICLR 2024.
  • Shah, Bikshandi, Zhang, Thakkar, Ramani and Dao (2024), FlashAttention-3 — arXiv:2407.08608.
  • Zadouri, Hoehnerbach, Shah, Liu, Thakkar and Dao (2026), FlashAttention-4 — arXiv:2603.05451, reporting roughly 1613 TFLOP/s on B200.

The version numbers track GPU generations, not algorithmic depth. The mathematics has not moved since 2022. Each release re-targets a new memory hierarchy and a new asynchronous instruction set, so read every FlashAttention benchmark as a statement about one chip.

What FlashAttention does not fix

Compute still scales as $\Theta(N^2 d)$. At long enough $N$ arithmetic dominates again, and only sparsity or a different operator helps — see sparse and block attention patterns and linear attention.

The KV cache is untouched. During decoding, memory is dominated by cached keys and values, which is the province of multi-query attention and multi-head latent attention.

Numerics are equivalent up to reassociation, not bitwise identical. Tests asserting exact equality against a reference implementation will fail, and should be written with a tolerance.

Beyond hand-written kernels

torch.nn.attention.flex_attention compiles a user-supplied score_mod or mask_mod into a fused tiled kernel, so a new masking scheme does not require new CUDA. It is documented as a prototype API as of PyTorch 2.13, so pin your version and expect signature churn. The FlashAttention-4 authors report that their CuTeDSL framework hosts FlexAttention-style and block-sparse variants without changes to the core. That is the direction of travel: one tiled skeleton, many masks.

What to learn next

What to learn next

These follow on from what you just read.

  • Fast Attention and Long Context

    Memory-bound vs compute-bound

    Most of what looks slow in a language model is not arithmetic, it is waiting for numbers to arrive from memory, and knowing which one you are in tells you what to fix.

  • Fast Attention and Long Context

    Multi-query attention

    Multi-query attention keeps many question-asking heads but gives them one shared set of keys and values, which shrinks the memory a model must re-read for every word it writes.

  • Fast Attention and Long Context

    Grouped-query attention

    Grouped-query attention gives every small group of attention heads one shared set of keys and values, which is why almost every model released since 2023 uses it.