Chapter 3 · Kernel optimization
Attention

Self-attention’s naive implementation is memory-bound for a specific, avoidable reason: computing softmax(QKT/√d)V requires materializing the full N×N score matrix, writing it out to , then reading it straight back for the softmax and the second matmul. None of that traffic does useful arithmetic, and an A100 has only 192KB of on-chip SRAM per streaming multiprocessor, across 108 of them, to avoid it with.

FlashAttention avoids ever writing that matrix to HBM. It loops over blocks of K and V in an outer loop, loading each block into SRAM once, then loops over blocks of Q in an inner loop, computing softmax incrementally across blocks, the same tiling idea from matrix multiplication applied to a running reduction instead of a running sum, and fusing every step of the attention computation into one kernel so intermediate results never leave the chip. The payoff is up to a 7.6× speedup on the attention computation itself for GPT-2, up to 9× fewer HBM accesses than the standard implementation, and, end to end, a 15% wall-clock speedup training BERT-large over the MLPerf 1.1 record and a 3× speedup on GPT-2.

IO-aware attention. One block of K and V is loaded into on-chip SRAM at a time while blocks of Q stream past it. The score block, the online softmax, and the multiply by V are fused into a single kernel, so the N×N score matrix is never written to HBM; a running max and sum rescale the partial output as each new block arrives. Shaded blocks mark one tile's worth of data.Illustrative numbers

The part of that trade which makes it the chapter’s thesis in miniature is easy to miss: FlashAttention does more floating-point arithmetic than the kernel it replaces. A backward pass needs the score matrix and the probability matrix softmax produced from it, and a standard implementation reads both back from HBM where the forward pass left them. FlashAttention never wrote them, so it stores the output and the softmax normalization statistics instead, the running maximum and running sum that the online softmax carries, and recomputes both matrices on chip from the blocks of Q, K, and V already sitting in SRAM. The paper calls that a form of selective gradient checkpointing and marks the difference from the ordinary kind: every implementation its authors knew of had to trade speed for memory, while this one, even with more FLOPs, speeds the backward pass up, because the arithmetic it adds is cheaper than the HBM accesses it removes.

The measurement is the argument. On an A100 running GPT-2 medium, FlashAttention has the higher FLOP count of the two implementations and the lower forward-plus-backward runtime, and HBM access is the primary factor affecting that runtime. Counting those accesses turns into a complexity result. With N the sequence length, d the head dimension, and M the size of SRAM, FlashAttention performs O(N2d2M-1) HBM accesses against Ω(Nd + N2) for the standard implementation, and the paper pairs that with a lower bound: no exact attention algorithm can asymptotically improve on the number of HBM accesses over all SRAM sizes. The tiling has a ceiling of its own, and the same sweep finds it. Widening the block shrinks HBM accesses and runtime together up to a block size of about 256, past which other factors such as arithmetic become the bottleneck and a larger block stops fitting in SRAM anyway.

Counting traffic rather than flops also rescues approximation. Approximate attention methods trade model quality for lower compute complexity and often do not achieve wall-clock speedup at all, because what they save is arithmetic and what the kernel is waiting on is memory. Block-sparse FlashAttention skips the blocks a block-sparsity mask zeroes, which multiplies the larger term of the IO complexity by s, the fraction of nonzero blocks, giving Θ(Nd + N2d2M-1s). That saving is in traffic, so it converts into time: 2 to 4× faster than FlashAttention itself, scaling up to sequence length 64K, and 2.8× faster than standard attention on Long-Range Arena while performing on par with it.

The original kernel still reached only 25 to 40% of a GPU’s theoretical peak FLOPs/s, and FlashAttention-2 traced the shortfall to work partitioning rather than the algorithm: for some problem shapes too few thread blocks were active at once, hurting , and shared-memory traffic between warps within a block was higher than it needed to be. Rebalancing that work, splitting the sequence dimension across more thread blocks and dividing work between warps to cut shared-memory communication, without changing the underlying algorithm, roughly doubled throughput over the original kernel, reaching 50 to 73% of theoretical peak FLOPs/s on an A100 and, in end-to-end GPT-style training, up to 225 TFLOP/s per A100, 72% model FLOPs utilization.

On Hopper the ceiling moved again: FlashAttention-2, tuned for the A100, achieves only 35% utilization on an H100, and FlashAttention-3 closes the gap with the same asynchronous machinery the previous section’s GEMM kernels use. Warp specialization overlaps computation with TMA data movement, block-wise matmuls are interleaved with the softmax so the tensor cores are not idle while exponentials run on slower units, and block quantization with incoherent processing exploits Hopper’s hardware FP8 support. The result is a 1.5 to 2.0× speedup on H100, up to 740 TFLOP/s at FP16, 75% utilization, close to 1.2 PFLOP/s at FP8, and 2.6× lower numerical error than a baseline FP8 attention.

All three kernels handle the case where a full sequence is processed at once, training or a single prefill pass. Decoding, where the KV cache grows by one token per step, needs the same IO-aware framing applied to a much smaller, much more frequent matmul, which is why inference engines run a separate, differently tuned attention kernel at decode time rather than reusing the training kernel unchanged.

In serving stacks that split has a name: FlashInfer, a library and kernel generator for inference, ships separate optimized kernels for prefill, decode, and mixed batching, behind unified attention, GEMM, and mixture-of-experts APIs that select among backends including FlashAttention-2 and 3, cuDNN, CUTLASS, and TensorRT-LLM. Its attention kernels operate directly on the paged and ragged KV-cache layouts the next chapters’ engines use, add cascade attention to share the KV cache of common prefixes across requests, and stay compatible with CUDA Graphs and torch.compile for low-latency serving.

Which FLOPs the hardware is fast at#

The first change on FlashAttention-2’s own list is not about parallelism at all. It is about which arithmetic the kernel performs, because on tensor-core hardware not all FLOPs cost the same. An A100 has a maximum theoretical throughput of 312 TFLOP/s of FP16/BF16 matmul and only 19.5 TFLOP/s of non-matmul FP32, which is to say each non-matmul FLOP is 16× more expensive than a matmul FLOP. Non-matmul FLOPs account for only a small fraction of attention’s total, and they still take longer to perform, so a kernel that wants more than half of theoretical peak has to spend as much of its time as possible inside the matmuls.

The non-matmul work in a fused attention kernel is the bookkeeping the online softmax demands, and FlashAttention-2 makes two tweaks to it that leave the output unchanged. The first is where the division goes. The original rescales both terms of the output update by the new normalizer every time a block arrives; FlashAttention-2 keeps an unscaled output accumulator plus the running sum alongside it, and divides by that sum once, after the last block. The second is what crosses to the backward pass: instead of saving both the running maximum and the running sum of exponentials, it stores only their combination, the logsumexp, the maximum plus the log of the sum. The rescaling by the ratio of successive maxima still happens per block, since that is what keeps the accumulator on a common scale, but the per-block division does not.

The second change is a loop inversion, and it is what makes the kernel parallel along the sequence. The original FlashAttention parallelizes over batch size and number of heads only, one thread block per attention head, which fills an A100’s 108 streaming multiprocessors well when that product is large, say at least 80. Long sequences are exactly the case where it is not: long context usually means a small batch, and the GPU runs out of thread blocks before it runs out of SMs. FlashAttention-2 swaps the loops, putting row blocks of Q on the outside and column blocks of K and V on the inside, which turns the outer loop over the sequence into an embarrassingly parallel one whose thread blocks never have to communicate. That is the occupancy the section above credits for the speedup, and the backward pass gets the same treatment along its own axis, one thread block per column block with atomic adds to accumulate the shared dQ update. The paper credits the swap and the sequence-dimension parallelism to Phil Tillet, who suggested and implemented both first in the Triton fused-attention tutorial.