Home › What Happens After You Hit Enter › FlashAttention (IO-aware attention kernels)
Pipeline stage · Operate

FlashAttention (IO-aware attention kernels)

Never write the full attention matrix to memory — compute it in tiles inside fast on-chip memory using a running softmax.

In one line

FlashAttention changes neither the mathematics nor the output; it just never writes the score matrix to memory, which is the entire reason long context became affordable.

Why you'd careThe thing you have already noticed

You ran a model that fits comfortably in VRAM, fed it an 8,000-token document, and got an out-of-memory error scaled like nothing you had allocated. Or you changed one config value — a different head dimension, a custom mask, fp32 instead of bf16 — and prefill got three times slower with nothing in the logs. Both are this stage. Textbook attention builds the full S-by-S score matrix in GPU memory: at 8,000 tokens and 32 heads in fp16 that is about 4 GB per layer, transient but real. FlashAttention computes the identical result in tiles that fit in on-chip SRAM and never writes that matrix down. When the kernel supports your configuration you get it. When it does not, you silently fall back to the version that does.

In and outWhat goes in, what comes out

InQ, K and V in HBM with shape [batch, heads, S, head_dim], a mask specification — causal, sliding-window or arbitrary — and tile sizes chosen so a Q block, a K block and a V block all fit in one multiprocessor's SRAM.
ProcessLoop over key blocks. Load a tile into SRAM, compute its partial scores, update a running row maximum and running sum, rescale the output accumulator to match the new maximum, add the block's contribution. Repeat until every key is consumed.
OutThe numerically exact attention output plus a per-row logsumexp statistic. HBM traffic scales with S rather than S squared, and the S-by-S score matrix is never written anywhere.

The result is exact algebraically and not bit-identical to the naive version, because the additions happen in a different order. What is lost is observability: there is no attention matrix left to inspect, so interpretability tooling, attention-based token attribution and any logging of attention weights all require a separate and much slower pass.

ConceptThe idea underneath

This is not a model change. Nothing about the network differs, and if the output distribution moves by more than floating-point noise, something is broken. It is a memory-hierarchy optimisation whose only ML content is one clever piece of numerics.

A GPU has two relevant tiers. HBM is large and comparatively slow — the FlashAttention paper measures an A100 at 40 GB of HBM at about 1.5 TB/s. On-chip SRAM is tiny and enormously faster: 192 KB per streaming multiprocessor at roughly 19 TB/s. Standard attention writes the score matrix and the softmax result to HBM and reads them back, moving O(S^2) bytes to produce only O(S * d) bytes of output. It is bandwidth-starved on data it never needed to keep.

The obstacle to tiling is softmax, which needs the maximum and the sum of an entire row before it can normalise anything, and a tile only sees part of the row. The fix is the online softmax: carry a running maximum and a running denominator, and whenever a new tile raises the maximum, rescale the accumulated output by the exponential of the difference before adding the new contribution. Every partial result stays correctable, so a row can be finished without ever having been assembled.

The payoff compounds with length. IO drops from quadratic to roughly linear in S, so memory stops limiting context and attention stops being bandwidth-starved. FlashAttention-2 reworked the partitioning of work across warps and thread blocks; FlashAttention-3 adds Hopper-specific asynchrony and low-precision paths. Same output, different arrangement of the same arithmetic.

At a glanceSee it

FlashAttention (IO-aware attention kernels) diagram

Attention computed tile by tile in SRAM with a running softmax; the score matrix is never stored.

The knobsHyperparameters and nuance

  • BLOCK_M and BLOCK_Nquery and key rows per tile, typically 64 or 128. Too large and the tile spills out of SRAM and occupancy collapses; too small and the amortisation disappears. Autotuned in Triton implementations, hard-coded in hand-written CUDA ones.
  • num_warps and num_stagesTriton scheduling parameters controlling parallelism and software-pipelining depth inside a tile. Mistuned values cost 20 to 50 percent throughput with no correctness effect at all.
  • attention backend selectionvLLM's --attention-backend flag (the VLLM_ATTENTION_BACKEND environment variable it replaced was deprecated in vLLM 0.13.0), or PyTorch's sdpa_kernel context manager. This is the knob that actually decides whether you get the fast kernel. Left on automatic, the dispatcher chooses, sometimes badly.
  • head_dimkernels enforce a maximum head dimension rather than a fixed menu of them. FlashAttention-1 topped out at 128; FlashAttention-2 and 3 support all head dimensions up to 256 in the forward pass, with tighter hardware requirements for the backward pass above 192. Exceeding a build's ceiling is a frequent cause of a silent fallback.
  • dtypethe fast kernels are bf16 and fp16 only. Running a model in fp32 out of caution disables FlashAttention entirely and no warning is emitted.
  • mask typecausal and sliding-window masks are supported natively; an arbitrary boolean mask forces a slower path or is not supported at all, which is a correctness problem rather than a speed one.

EffectHow this stage moves the answer

By design, nothing. The output distribution is what the textbook implementation would give, and a quality change after switching kernels should be treated as a masking bug rather than a tradeoff. Two second-order effects do reach the user. The first is reachability: without IO-aware kernels a 100,000-token prompt is not slow, it is impossible, so every answer that depends on the model seeing a whole codebase or a whole contract exists because of this stage. The second is divergence. Reordered accumulation shifts logits in the last few decimal places, and when the top two candidates are nearly tied that is enough to flip the choice, after which the completions separate entirely. Same weights, same prompt, temperature 0, different kernel, visibly different answer. It is not a quality difference, but it will look like one in a diff.

EvalsWhat it does to your measurements

Expect a flat line on quality and a large move on prefill latency and maximum feasible context. That expectation is itself a test: run a fixed prompt set through two backends and compare token-level agreement. A disagreement rate above what floating-point noise explains points at a mask or head-dimension bug, not at an optimisation. The silent fallback is what corrupts benchmark runs. A configuration change that drops you onto the naive path produces a two-to-five-times prefill regression with no error, so the number regresses and gets attributed to the model, the batch size or the scheduler. Log the selected backend on every run and treat it as part of the result. Never compare prefill numbers across framework versions without checking it either — kernel dispatch heuristics change between releases far more often than weights do.

Failure modesWhen it goes wrong

  • Out-of-memory during prefill on a model whose weights fit easilythe naive path materialising an S-by-S score matrix per head, which grows quadratically while nothing you allocated does.
  • Prefill two to five times slower after a config change, nothing in the logsan unsupported head dimension, dtype or mask silently dispatched to the fallback kernel.
  • Wrong answers with a custom attention mask that works on the eager paththe fused kernel supports only causal and windowed masks and ignores or mishandles arbitrary ones.
  • NaNs appearing only at long sequence lengthsa fully masked row leaves the running maximum at negative infinity and the denominator at zero; correct kernels guard this, locally patched ones often do not.
  • Outputs differ between two machines running identical codedifferent GPU architectures select different kernels and different tile sizes, changing accumulation order.

PapersWhere this comes from

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-AwarenessDao et al., 2022. arXiv:2205.14135. Established that attention is IO-bound rather than compute-bound and gave the tiled, recomputation-based exact algorithm that made long context practical.
  • FlashAttention-2: Faster Attention with Better Parallelism and Work PartitioningDao, 2023. arXiv:2307.08691. Reworked the partitioning of work across thread blocks and warps for roughly a further 2x, and is the version most serving stacks actually ship.
  • Online normalizer calculation for softmaxMilakov & Gimelshein, 2018. arXiv:1805.02867. The single-pass running-maximum softmax that makes tiled attention possible at all; without it a tile cannot normalise anything.
  • Self-attention Does Not Need O(n^2) MemoryRabe & Staats, 2021. arXiv:2112.05682. Showed independently that exact attention is computable in constant memory by chunking, framing the memory claim that FlashAttention then made fast on real hardware.
A living map of modern AI — kept current every morning