TOPIC #211Advanced 14 min read

Flash Attention

AI
AI & ML Editorial
Report an issue
Key takeawayCore Concept Summary

The IO-aware exact attention kernel that made long context affordable: attention is reorganized so the N² score matrix never travels to slow GPU memory — tiles are processed on-chip with an online softmax. Plus the FA1 → FA2 → FA3 evolution on Ampere, Hopper, and Blackwell GPUs.

Flash Attention: Tiled Exact Attention Keeping S Off HBM ⚡

Q, K, V are processed in blocks; softmax is computed online (rescaling running statistics) so the N×N score matrix is never materialized in HBM — attention becomes IO-bound instead of memory-wasted.

Flash Attention: Tiled Exact Attention Keeping S Off HBM ⚡
100%
Touchpad: Pinch to zoom • Drag to pan
Rendering visual architecture flowchart...

01.The Problem: The Warehouse Trip Is the Expensive Part

To compute attention naively, a GPU does two matmuls with a softmax in between:

S = softmax(QKᵀ / √d), then Out = S·V.

The problem is the middle object: S is an N×N score matrix per head, per layer — and the standard implementation writes it all the way out to the GPU's big slow memory (HBM), then reads it back for the second matmul, then writes gradients back over it.

So ask

Insight

Is the GPU actually busy doing math while that giant matrix ferries back and forth?

Mostly no. An A100's tensor cores can do the matmuls quickly, but HBM moves only ~2 TB/s — the chip spends its time hauling data, not computing. The traffic is Θ(N²) reads and writes, and it dominates wall time.

Two memories, two speeds — that is the whole story:

  • HBM (GPU DRAM): huge (40-80 GB) but slow-ish, ~3-8 TB/s. The warehouse across the street.
  • SRAM (on-chip): tiny (~tens of MB) but ~19 TB/s — the workbench you are standing at.

Flash Attention (Dao et al., 2022, arXiv 2205.14135) reframed attention as an IO problem: the enemy is HBM round-trips, not FLOPs.

02.The Idea in Plain Words: Never Let the Giant Matrix Exist in Memory

Flash Attention in one line

Insight

Compute attention in small tiles on the fast on-chip workbench, fusing matmul → softmax → matmul, so the N² score matrix is never written to slow memory at all.

The kernel guarantees HBM is touched only for things you actually needed anyway:

  • in: Q, K, V
  • out: O
  • plus two tiny per-row statistics: the running max m and running denominator ℓ of softmax

The N² intermediates live only in registers/SRAM, tile by tile, and die there.

Consequences, all three of them:

  • Memory: Θ(N²) → Θ(N) — attention becomes memory-linear, enabling longer sequences and bigger batches.
  • Exactness: this is NOT an approximation. The output is mathematically identical to naive attention (up to floating-point reassociation order). Nothing is dropped, unlike sparse or linear attention.
  • Speed: 2-4× end-to-end attention speedups at realistic sequence lengths; 7.6× on the attention kernel alone in the FA1 benchmarks.

How can you softmax a row without ever having the whole row? The trick that makes it work — online softmax — is the next section.

03.A Simple Worked Example: Softmax Over Two Small Blocks

Softmax needs two things about a whole row: its maximum (for numerical stability) and the sum of exponentials. Normally you need the full row. Flash Attention keeps running versions and corrects them as new tiles stream past.

Row of scores: [1, 2, 5], streamed as block A = [1, 2] then block B = [5].

After block A:

  • running max m = 2
  • running sum ℓ = e^{1−2} + e^{2−2} = 0.368 + 1 = 1.368
  • partial weights for entries 1, 2 (relative to A alone): 0.368/1.368, 1/1.368

Block B arrives and contains a BIGGER number than the current max. Correction time:

  • new max m = 5
  • every previously accumulated value was computed with max 2, so scale it by e^{m_old − m_new} = e^{−3} = 0.0498:
  • ℓ = 1.368 × 0.0498 + e^{5−5} = 0.068 + 1 = 1.068

Final weights: [0.368×0.0498, 1×0.0498, 1] / 1.068 = [0.017, 0.047, 0.936]

Now compute the ordinary softmax of [1, 2, 5] in one shot: e¹+e²+e⁵ = 2.718+7.389+148.41 = 158.5 → weights [0.017, 0.047, 0.936]. Identical. The streaming version just deferred the bookkeeping.

The same rescaling also multiplies the running output accumulator O (the partial S·V sums), so the weighted values stay consistent even though blocks were processed in arbitrary order.

04.Visual Intuition: Crates, Not the Whole Truck

code
   HBM (warehouse, slow truck)          SRAM (workbench, fast)
 ┌───────────────────────────┐      ┌─────────────────────────────┐
 │ Q ──► fetch ONE row block  │      │  q-tile in registers        │
 │ K ──► stream ONE col block │ ───► │  s-tile = q·kᵀ  (small!)    │
 │ V ──► same col block       │      │  update m, ℓ, rescale O     │
 │                            │      │  discard s-tile ✗           │
 │ O ◄── write final rows only│ ◄─── │  next block…                │
 └───────────────────────────┘      └─────────────────────────────┘
     the ONLY things crossing          the N² grid exists here,
     the aisle: Q, K, V, O, m, ℓ        briefly, one tile at a time

The naive version, for contrast:

code
  compute full N×N S ──► WRITE to warehouse ──► READ all back
                       ──► compute P=softmax ──► WRITE again
                       ──► READ for P·V          (Θ(N²) aisle trips!)

Flash Attention simply never makes those trips. Total HBM traffic falls from Θ(N²) to Θ(N) — the memory footprint becomes linear in sequence length, which is also exactly why 32K-token rows (2.1 GB per head as a materialized matrix) stop being a problem.

05.The Analogy: The Short-Order Cook with a Tiny Counter

Carry one picture through: a cook catering 10,000 bowls of soup, with one small counter and a warehouse across the street.

  • The naive method: haul every ingredient to the warehouse-adjacent mega-table, lay out all N² taste-combinations, record them, carry them back to season each bowl. The kitchen spends all day on hauling, and the mega-table does not even fit.
  • Flash Attention's cook: works one crate at a time on the fast counter. Mix a tile, taste it, fold the result into the pot.
  • The running max and sum (m, ℓ) = the notepad beside the pot: "hottest pepper so far: 5; running flavor total: 1.068". When a crate turns out hotter than anything before, the cook does NOT re-taste the whole warehouse — just rescales the notes (× e^{old−new}) and continues. That is online softmax, literally: an accountant of exponentials who never needs the full ledger at once.
  • Backward pass (training): instead of keeping every tasted cup labeled (storing the N² matrix), the cook writes down only the recipe inputs and the notepad, and re-tastes specific batches when the inspector asks what went in (recompute tiles from Q, K, V, O, m, ℓ). Re-tasting is cheap compared to warehousing 10,000 cups.

Everything the rest of this topic covers — tiling orders, warp splits, asynchronous copies — is this cook optimizing the choreography between counter, notepad, and warehouse trips.

06.Inside the Kernel: Tiling, Recomputation, and FA1 → FA2 → FA3

Two design decisions do the heavy lifting:

Streaming blocks with online softmax. Process K/V column-blocks one at a time, maintaining per-row running max mᵢ and running denominator ℓᵢ. When a new block brings a larger max, previously accumulated outputs are rescaled by exp(m_old − m_new) — the same algebra as numerically-stable softmax, just deferred.

Backward recomputation. The backward pass has no saved S/P to read, so Flash Attention recomputes tiles from (O, m, ℓ) plus Q, K, V. That is a deliberate FLOPs-for-IO trade: extra compute is nearly free compared with HBM round-trips.

The three-generation evolution:

  • FA1 (2022): also used warps split along the row dimension with careful shared-memory synchronization.
  • FA2 (Tri Dao, 2023, arXiv 2307.08691) improved it by:
    1. Splitting work across sequence length, not just heads/batch — removing the inter-block synchronization overhead (warp partitioning over the N dimension instead of the key/row partitioning of FA1).
    2. Interleaving matmul and softmax ops, moving rescaling out of the inner loop and deferring it to the epilogue.
    3. Parallelism across KV blocks with atomic-free split-KV for long sequences and decode. Net effect: ~2× faster than FA1, ~50-75% of theoretical peak FLOPs on A100 — and roughly 2× better utilization than Triton reference kernels that made the same asymptotics easy to reproduce.

07.Flash Attention in the 2024-2026 Stack

FlashAttention became infrastructure, not a paper: PyTorch's SDPA backend, vLLM, TensorRT-LLM, and virtually every training framework default to FA2/FA3-class kernels (or cuDNN's fused attention, which matched FA3-class performance around 2025). Decode-time variants — Flash-Decoding (parallel split over the KV cache) and paged kernels like FlashInfer — extend the tiling idea to the growing-cache regime of generation.

Where are its limits? The cook analogy answers cleanly: FA removes the hauling (memory traffic and constants) but the number of taste-combinations is still N² — it does not change quadratic FLOPs. That is why the companion sparse-attention topic (NSA/DSA — never tasting most crate pairs at all) and KV-cache work (PagedAttention) sit next to it, not instead of it. FA is the exact, dense, IO-optimal floor every later technique is measured against.

python— Using FlashAttention as a drop-in for PyTorch scaled_dot_product_attention
from flash_attn import flash_attn_func  # FA2 (FA3: flashinfer / hopper kernels)

q = kv_proj_q(x)  # (batch, seqlen, n_heads, head_dim) - no transpose needed
k = kv_proj_k(x)
v = kv_proj_v(x)

# Exact attention, Θ(N) HBM footprint, causal masking fused
out = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)

# Memory saving vs naive: at N=32k, per-head score matrix is
# 32_768² × 2 bytes ≈ 2.1 GB — never allocated with FA.

Architectural Trade-offs & Production Realities

Architectural Advantages

  • Exact attention with Θ(N) memory instead of Θ(N²) — long-context training/serving becomes possible.
  • 2-4× end-to-end speedups; FA2 hits 50-75% of A100 peak; FA3 exploits Hopper asynchrony to ~75%/1.2 PFLOPS FP8.
  • Drop-in fusion: no model-quality compromise (unlike approximate linear attention).

Trade-offs & Constraints

  • Does not fix the O(N²) FLOPs — quadratic compute at 1M context still requires sparse methods.
  • Kernel complexity: per-architecture tuning (Ampere/Hopper/Blackwell); head-dim and masking variants need re-derivation (e.g., softcap, alibi).
  • Backward trades compute for IO via recomputation — extra matmuls in the backward pass.
Production Implementation in Big Tech
vLLM / PyTorch ecosystem• Default attention backend for production LLM inference

vLLM ships FlashAttention (and FlashInfer) kernels behind its attention layers, combining FA's Θ(N) IO with PagedAttention block tables; PyTorch 2.x exposes torch.nn.functional.scaled_dot_product_attention with an FA backend, so millions of training jobs and API deployments run fused exact attention by default rather than naive einsum softmax.

Staff+ Engineering Takeaways

  • Flash Attention is exact attention reorganized as an IO problem: N² score matrices never touch HBM.
  • Online softmax with running max/denominator lets blocks of K/V stream through SRAM; backward recomputes tiles to save memory.
  • FA2 (2023) ~2× FA1 via sequence-level parallelism and fewer non-GEMM ops; FA3 (2024) ~1.5-2.5× FA2 on Hopper with TMA + warp specialization + FP8.
  • Memory drops Θ(N²) → Θ(N), enabling long-context training and decode kernels like Flash-Decoding/FlashInfer.
  • FA does not remove quadratic FLOPs — it is the dense floor that sparse attention builds on top of.

Topic Knowledge Check

Exercise 1 of 3 • Test your architectural comprehension.

Exercise 1 of 30 answered
1

What is the fundamental resource Flash Attention optimizes?

Rate This Architecture ChapterFeedback & Rating

How clear and actionable was this distributed systems breakdown?