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 N×NN \times N attention matrix S and the softmax output P in HBM — reading and writing O(N2)O(N^2) elements for a sequence of length NN. When NN 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.

Open in Lab
Explore the GPU memory hierarchy. Click each level to see its capacity and bandwidth. Notice the 10× bandwidth gap between HBM and SRAM.
The demo wakes as you arrive…

FlashAttention recap: tiling and online softmax

Before diving into the improvements, let us recall what FlashAttention does. Standard attention computes S=QK⊤\text{S} = QK^\top, applies softmax row-wise to get PP, then multiplies O=PVO = PV. This requires materializing the full N×NN \times N 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 mm and a running sum of exponentials ℓ\ell. 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 O(N2)O(N^2) to O(N)O(N) and achieves 2-4× speedup by minimizing HBM reads and writes. But it still leaves performance on the table.

Open in Lab
Watch how tiling processes the attention matrix block by block. Each block is loaded into SRAM, computed, and accumulated — the full N×N matrix never materializes.
The demo wakes as you arrive…

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 diag(ℓ(j))−1\text{diag}(\ell^{(j)})^{-1}:

O(j)=diag(ℓ(j−1)/ℓ(j))⋅O(j−1)+diag(ℓ(j))−1⋅eS(j)−m(j)V(j)O^{(j)} = \text{diag}(\ell^{(j-1)} / \ell^{(j)}) \cdot O^{(j-1)} + \text{diag}(\ell^{(j)})^{-1} \cdot e^{S^{(j)} - m^{(j)}} V^{(j)}

FlashAttention-2 observes that this rescaling can be deferred. Instead of dividing by ℓ\ell at every step, the algorithm maintains an unscaled running output O~\tilde{O} and only divides once at the very end:

O~(j)=diag(em(j−1)−m(j))⋅O~(j−1)+eS(j)−m(j)V(j)\tilde{O}^{(j)} = \text{diag}(e^{m^{(j-1)} - m^{(j)}}) \cdot \tilde{O}^{(j-1)} + e^{S^{(j)} - m^{(j)}} V^{(j)}

This eliminates one division per block — a non-matmul operation. Additionally, instead of storing both the running max mm and sum ℓ\ell for the , FlashAttention-2 stores only their combination as the logsumexp L=m+log⁡(ℓ)L = m + \log(\ell). This halves the bookkeeping overhead and simplifies the backward pass.

O~(j)=diag(em(j−1)−m(j)) O~(j−1)+eS(j)−m(j)V(j),O=diag(ℓ)−1 O~(last)\tilde{O}^{(j)} = \text{diag}(e^{m^{(j-1)} - m^{(j)}}) \, \tilde{O}^{(j-1)} + e^{S^{(j)} - m^{(j)}} V^{(j)}, \qquad O = \text{diag}(\ell)^{-1}\,\tilde{O}^{(\text{last})}
Deferred rescaling — divide only once at the end — Instead of normalizing by the softmax denominator ℓ\ell at every block, the algorithm accumulates unnormalized products and applies a single normalization at the end. Each block only needs an exponential rescaling to match running max values — a cheaper operation — while the expensive division by ℓ\ell happens just once. This reduces non-matmul FLOPs without changing the final output.

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.

Open in Lab
Compare how thread blocks are scheduled. FlashAttention uses batch × heads. FlashAttention-2 adds the sequence dimension, filling many more SMs.
The demo wakes as you arrive…

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 QK⊤QK^\top, 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 QK⊤QK^\top 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.

Open in Lab
Compare split-K (FlashAttention) vs split-Q (FlashAttention-2). Notice how split-Q eliminates the shared memory synchronization step entirely.
The demo wakes as you arrive…

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 −∞-\infty. 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 O(N)O(N) memory instead of O(N2)O(N^2).

Open in Lab
Click the bars to compare throughput across implementations. FlashAttention-2 consistently achieves ~2× the throughput of FlashAttention.
The demo wakes as you arrive…

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.

FlashAttention-2 Forward Pass (Pseudocode)python

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

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

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

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

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

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