Systems & Optimization2023advanced13 min read
FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
FlashAttention-2: انتباه أسرع بتوازٍ أفضل وتوزيع عمل أذكى
Dao, Tri — ICLR 2024
The problem
brought from quadratic memory to linear and achieved 2-4× speedup, but it still reached only 25-40% of the 's theoretical maximum . Optimized matrix multiplication (GEMM) routinely hits 80-90%. The gap came from suboptimal work partitioning between thread blocks and warps, plus unnecessary non-matmul floating-point operations that underutilized the GPU's specialized .
The contribution
Three targeted optimizations that together double FlashAttention's speed. First, algorithmic tweaks that reduce non-matmul by deferring rescaling to the end and storing only the logsumexp instead of separate max and sum. Second, parallelizing the outer loop over the dimension so that long sequences with small batch sizes still saturate all streaming multiprocessors. Third, switching from a split-K to a split-Q partitioning strategy within each , eliminating inter-warp shared memory synchronization in the .
The impact
FlashAttention-2 became the de facto attention kernel for and serving large language models. By reaching 50-73% of theoretical GPU throughput, it made long-context training (16k+ tokens) economically viable. It is integrated into PyTorch, Hugging Face, and virtually every major LLM training framework. The ideas directly influenced FlashAttention-3, Flash-Decoding, and the broader push toward hardware-aware algorithm design in deep learning.
Imagine a restaurant kitchen during the dinner rush. The original FlashAttention was like a brilliant chef who learned to prep ingredients in small batches so the tiny counter isn't overwhelmed — no more hauling everything to the walk-in fridge and back.
But the kitchen was still slow because the sous-chefs kept bumping into each other, passing plates through a single window, and waiting for each other to finish.
FlashAttention-2 redesigns who does what: each sous-chef gets their own station, works on their own dishes, and never needs to wait for a hand-off. The recipes are the same, the food is identical — but the kitchen now runs at nearly full speed.
The hardware reality: why memory matters more than math
Modern GPUs are not just fast calculators — they are complex systems with a . At the top sits (High Bandwidth Memory): large (40-80 GB on an A100) but relatively slow (1.5-2.0 TB/s bandwidth). On the chip itself is (on-chip shared memory): tiny (192 KB per ) but blazingly fast (~19 TB/s bandwidth). This is a 10× speed gap.
The key insight behind both FlashAttention and FlashAttention-2 is that attention is memory-bound, not compute-bound. The GPU spends more time moving data between HBM and SRAM than it does actually computing. Standard attention materializes the full attention matrix S and the softmax output P in HBM — reading and writing elements for a sequence of length . When is 4096 or more, this dominates the wall-clock time.
GPUs also have specialized Tensor Cores designed for matrix multiplication. On an A100, Tensor Cores deliver 312 TFLOPs/s for FP16/BF16 matmul, but general-purpose FP32 arithmetic runs at only 19.5 TFLOPs/s. That means each non-matmul FLOP is effectively 16 times more expensive than a matmul FLOP. This asymmetry is central to FlashAttention-2's design: spend as much time as possible on matmul and minimize everything else.
FlashAttention recap: tiling and online softmax
Before diving into the improvements, let us recall what FlashAttention does. Standard attention computes , applies softmax row-wise to get , then multiplies . This requires materializing the full matrices S and P in HBM.
FlashAttention avoids this by : it divides Q, K, V into small blocks that fit in SRAM. For each pair of blocks, it computes the local attention scores, applies a local softmax, multiplies with the corresponding V block, and accumulates the result. The key challenge is that softmax is a global operation — it needs the maximum and sum over the entire row. solves this by maintaining running statistics: a running max and a running sum of exponentials . Each time a new block arrives, the statistics are updated, and the partial output is rescaled so the final result is exact — no approximation.
This tiling approach reduces memory from to and achieves 2-4× speedup by minimizing HBM reads and writes. But it still leaves performance on the table.
Improvement 1: fewer non-matmul operations
The first optimization targets the online softmax rescaling. In the original FlashAttention, after processing each new K-V block, the algorithm rescales both terms of the output update by :
FlashAttention-2 observes that this rescaling can be deferred. Instead of dividing by at every step, the algorithm maintains an unscaled running output and only divides once at the very end:
This eliminates one division per block — a non-matmul operation. Additionally, instead of storing both the running max and sum for the , FlashAttention-2 stores only their combination as the logsumexp . This halves the bookkeeping overhead and simplifies the backward pass.
Improvement 2: parallelism along the sequence dimension
The original FlashAttention parallelizes only over the and number of attention heads. Each thread block handles one head of one sequence. On an A100 with 108 streaming multiprocessors (SMs), you need at least 80 or so thread blocks to keep the GPU busy. With a batch size of 1 and 32 heads, that is only 32 thread blocks — many SMs sit idle.
FlashAttention-2 adds a new dimension of : the sequence length. The outer loop iterates over row blocks of Q, and each iteration is independent (it reads from all of K and V but writes to its own slice of the output). This is embarrassingly parallel — no communication needed between row blocks.
With this change, the total number of thread blocks becomes batch × heads × (N / block_size), which for long sequences can be thousands. Now even a single long sequence with one head can fully occupy the GPU. This is exactly the regime that matters most — long-context training with large models and small batch sizes.
In the backward pass, parallelism works similarly but with a twist. The outer loop iterates over column blocks of K and V (not row blocks of Q). Each column block accumulates dK and dV independently. The update to dQ, however, is shared across column blocks — multiple thread blocks may need to add to the same dQ slice. This is handled with atomic additions, a lightweight synchronization primitive that allows concurrent updates without full barriers.
Improvement 3: smarter work partitioning within warps
Even within a single thread block, there are multiple warps (groups of 32 threads) that need to divide the work. How they divide it determines how much they need to communicate through shared memory — and shared memory communication is slow.
In FlashAttention, the strategy was split-K: K and V are split across 4 warps, while all warps share Q. Each warp computes a slice of , then multiplies with its slice of V. But now the partial results must be combined — each warp writes its piece to shared memory, all warps synchronize, then one warp reads and sums them. This write-sync-read cycle is a bottleneck.
FlashAttention-2 flips the strategy to split-Q: Q is split across warps, while all warps share K and V. Now each warp computes its own slice of using the full K, multiplies with the full V, and directly produces its slice of the output. No communication between warps is needed at all. Think of it as giving each worker their own question list (Q) but letting everyone read from the same reference book (K, V) — each worker produces their own answers independently.
Bonus optimization: smarter causal masking
In language modeling, a ensures each token can only attend to previous tokens. In the attention matrix, this means every entry above the diagonal is . Since FlashAttention already works in blocks, any block where all column indices exceed the row indices can be skipped entirely — roughly half the blocks for long sequences. This alone gives 1.7-1.8× speedup.
FlashAttention-2 refines this further: for blocks that straddle the diagonal, the mask only needs to be applied within the one block that crosses it. All blocks fully below the diagonal need no masking at all. This reduces the masking overhead to a single block per row, minimizing the cost of what would otherwise be an element-wise operation applied to many elements.
Results: doubling the throughput
The combined effect of the three improvements is substantial. On an A100 80GB GPU, FlashAttention-2 achieves up to 230 TFLOPs/s in the forward pass (73% of theoretical max) and up to 196 TFLOPs/s in the backward pass (63% of max). For comparison, FlashAttention reaches about 124 TFLOPs/s forward and 113 TFLOPs/s backward. That is approximately a 2× speedup across the board.
Compared to standard PyTorch attention, the gap is even larger: FlashAttention-2 is 3-10× faster depending on sequence length and head dimension, while simultaneously using memory instead of .
In end-to-end training, FlashAttention-2 reaches 225 TFLOPs/s per A100 GPU when training GPT-style models with 2.7 billion parameters on 8k context. This represents 72% model FLOPs utilization — remarkably close to the hardware limit. For comparison, the same training without any FlashAttention achieves only 80 TFLOPs/s, a 2.8× slowdown. The improvement is most dramatic for long sequences: at 8k context, the speedup over baseline is 2.8×, while at 2k context it is 1.4×.
Multi-query and Grouped Query Attention support
Modern LLMs increasingly use multi-query attention (MQA) and (GQA) to reduce the size of the KV cache during . In MQA, all query heads share a single key-value head. In GQA, groups of query heads share a key-value head. Both reduce memory bandwidth requirements during decoding.
FlashAttention-2 handles these variants natively: instead of physically duplicating K and V heads to match the number of Q heads, it manipulates index pointers so the same K-V data is read by multiple Q groups. In the backward pass, the gradients dK and dV are summed across the query heads that share them. This support is important because it means FlashAttention-2 accelerates not just training but also the inference architectures used in production.
The full FlashAttention-2 forward pass
Putting it all together, here is the complete forward algorithm in pseudocode. The outer loop iterates over row blocks of Q (each assigned to a separate thread block for parallelism). The inner loop iterates over column blocks of K and V. At each step, the algorithm computes local attention scores, updates running softmax statistics, and accumulates the unscaled output. Only at the end is the output normalized.
Simplified to show the idea — not the real implementation.
# Q, K, V: [N, d] in HBM; Br, Bc: block sizes
# Outer loop: each iteration = 1 thread block (parallel over rows) for i in range(ceil(N / Br)):
Qi = load_from_HBM(Q[i*Br : (i+1)*Br]) # Load Q block to SRAM
Oi = zeros(Br, d) # Accumulator (unscaled)
li = zeros(Br) # Running sum of exponentials
mi = full(Br, -inf) # Running row-wise max
# Inner loop: iterate over K, V blocks
for j in range(ceil(N / Bc)):
Kj, Vj = load_from_HBM(K[j], V[j]) # Load K, V block
Sij = Qi @ Kj.T # Local attention scores
# Update running softmax statistics
mi_new = max(mi, rowmax(Sij))
P_tilde = exp(Sij - mi_new) # Local softmax numerator
li = exp(mi - mi_new) * li + rowsum(P_tilde)
# Accumulate unscaled output (NO division by li here!)
Oi = diag(exp(mi - mi_new)) * Oi + P_tilde @ Vj
mi = mi_new
# Single normalization at the very end
Oi = diag(1 / li) * Oi
Li = mi + log(li) # Store logsumexp for backward pass
write_to_HBM(Oi, Li)Timeline: the FlashAttention lineage
2018
Online Softmax (Milakov & Gimelshein)
Proposed computing softmax in a single pass by maintaining running max and sum statistics. This foundational technique enables block-wise attention without materializing the full attention matrix.
2022
FlashAttention (Dao et al.)
Combined tiling, online softmax, and recomputation into an IO-aware attention algorithm. Reduced memory from O(N²) to O(N) and achieved 2-4× speedup with no approximation. Widely adopted across the industry.
2023
FlashAttention-2 (this paper)
Doubled FlashAttention's speed by reducing non-matmul FLOPs, parallelizing over sequence length, and switching to split-Q warp partitioning. Reached 73% of theoretical max throughput on A100 GPUs.
2023
Flash-Decoding (Dao et al.)
Extended FlashAttention-2 ideas to inference by parallelizing over the KV cache length dimension, enabling efficient long-context decoding.
2024
FlashAttention-3 (Shah et al.)
Leveraged Hopper GPU features (TMA, 4th-gen Tensor Cores, FP8) to push attention efficiency even closer to hardware limits on H100 GPUs.
CitationTri Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. ICLR, 2024.
Terms in this paper
- FlashAttentionFlashAttention
- Attentionآلية الانتباه
- Tilingالتبليط
- Online Softmaxsoftmax المتدفق
- SRAMالذاكرة الساكنة
- HBMالذاكرة عالية النطاق
- GPUوحدة معالجة الرسوميات
- Kernel Fusionدمج النواة
- Throughputمعدل التدفق والإنتاجية
- Forward Passالتمرير الأمامي
- Backward Passالتمرير الخلفي
- Causal Maskقناع سببي
- Multi-Head Attentionالانتباه المتعدد المسارات
- Grouped Query Attentionانتباه الاستعلام المُجمَّع
- Softmaxسوفت ماكس