Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

15. FlashAttention

In this chapter

  • Why standard attention is limited by memory traffic, not arithmetic.
  • The online-softmax recurrence, derived step by step: how to normalize scores you haven't all seen yet.
  • FlashAttention as tiling plus the online recurrence, in PyTorch, C++, Rust and Triton.
  • Flash-decoding: splitting a long KV cache across many blocks and merging the partial results.

You will build

online_attention in engine/attention.py and a FlashAttention forward kernel in engine/kernels/triton_flash.py, with causal masking, query offsets and grouped-query heads.

Time: 6-8 hours. GPU: optional.

The score matrix is the problem

Standard attention (Chapter 5) computes $S = QK^\top/\sqrt d$, then $P = \operatorname{softmax}(S)$, then $O = PV$. For a sequence of $T$ tokens, $S$ and $P$ are $T \times T$ per head. At $T = 8{,}192$, each is 67 million entries: 256 MiB per head in FP32. With 32 heads and a batch of 8, the intermediate scores alone would need 64 GiB.

Even when the memory fits, the traffic kills performance. Unfused, $S$ is written to HBM, read back for the softmax, $P$ is written, then read again for $PV$: four passes over $T^2$ numbers. All that just to produce an output of size $T \times d$, with $d$ typically 64-256. By Chapter 10’s roofline, unfused attention is memory-bound at long sequence lengths even though its arithmetic is large.

FlashAttention (Dao et al., 2022) computes exactly the same output without ever storing $S$ or $P$ in global memory. It tiles $Q$, $K$ and $V$ into on-chip memory and uses a clever recurrence for the softmax. The output matches standard attention up to rounding; it isn’t an approximation. Sparse and linear attention (Chapters 28-29) are approximations, or different models altogether. This isn’t.

The obstacle: softmax needs the whole row

Tiling a matmul works because a dot product is a sum: process $K$ in chunks and add up partial sums. Softmax is harder. Each weight is $e^{s_j}/\sum_k e^{s_k}$, and the denominator needs every score in the row. The stable version also subtracts the row’s maximum, which you don’t know until you’ve seen every score. It looks like you need the whole row first.

Online softmax: normalize as you go

The trick (Milakov and Gimelshein, 2018) is to keep a running answer that’s correct for the scores seen so far, and to fix it up when new scores arrive. For one query row, keep three running values:

  • $m$: the largest score seen so far,
  • $\ell$: the sum of $e^{s_j - m}$ over scores seen so far,
  • $a$: the sum of $e^{s_j - m}, v_j$ over scores seen so far (a vector of length $d$, not yet normalized).

When a new tile of scores arrives with a larger maximum $m_{\text{new}}$, every old term $e^{s_j - m_{\text{old}}}$ must become $e^{s_j - m_{\text{new}}}$. Since

$$ e^{s_j - m_{\text{new}}} = e^{s_j - m_{\text{old}}} \cdot e^{m_{\text{old}} - m_{\text{new}}}, $$

every old term is fixed by multiplying by the same factor $\alpha = e^{m_{\text{old}} - m_{\text{new}}}$. So the update for a tile with scores $s$ and values $V_{\text{tile}}$ is:

$$ \begin{aligned} m_{\text{new}} &= \max!\big(m_{\text{old}},, \max(s)\big) \ p &= e^{,s - m_{\text{new}}} \ \ell_{\text{new}} &= \alpha,\ell_{\text{old}} + \textstyle\sum p \ a_{\text{new}} &= \alpha, a_{\text{old}} + p, V_{\text{tile}} \end{aligned} \qquad\text{and at the end}\qquad O = a / \ell . $$

Both $\ell$ and $a$ must be rescaled. Forgetting one mixes terms measured against different references, a classic bug.

Worked example

One query, scores $[0, 1, 2, -1]$ against scalar values $[2, 4, 1, 3]$, processed in two tiles of two.

Tile 1, scores $[0, 1]$: $m = 1$, $p = [e^{-1}, e^0] = [0.3679, 1]$, $\ell = 1.3679$, $a = 0.3679\cdot 2 + 1\cdot 4 = 4.7358$.

Tile 2, scores $[2, -1]$: the maximum rises to $m = 2$, so $\alpha = e^{1-2} = 0.3679$. $p = [e^0, e^{-3}] = [1, 0.0498]$.

$$ \ell = 0.3679 \times 1.3679 + 1 + 0.0498 = 1.5530, \qquad a = 0.3679 \times 4.7358 + 1\cdot 1 + 0.0498 \cdot 3 = 2.8915 . $$

$O = 2.8915 / 1.5530 = 1.8619$, exactly what softmax([0,1,2,-1]) @ [2,4,1,3] gives. Your milestone tests check this example.

Explore: the online softmax, one tile at a time

Edit the scores and values and step through the tiles. Watch the running maximum jump and the old sums get rescaled; the final output always matches the dense softmax.

The algorithm in plain PyTorch

Here’s the recurrence over key tiles, for all queries and heads at once, with causal masking by position:

def online_attention(q, k, v, query_positions=None, key_positions=None, tile=16):
    """The same result as causal_attention, computed one key tile at a time without ever
    holding the full [T, S] score matrix. This is FlashAttention's recurrence in plain
    PyTorch.  (Your engine: Chapter 15)

    Running state per query row: m (max score so far), l (sum of exp(score - m)),
    acc (sum of exp(score - m) * value). A new tile with a larger maximum rescales the old
    l and acc by exp(m_old - m_new) so every term shares one exponent reference.
    """
    if tile < 1:
        raise ValueError("tile must be positive")
    batch, q_heads, t, d = q.shape
    kv_heads, s = k.shape[1], k.shape[2]
    if q_heads % kv_heads:
        raise ValueError("Query heads must be a multiple of KV heads")
    group = q_heads // kv_heads
    k = k.repeat_interleave(group, dim=1) if group > 1 else k
    v = v.repeat_interleave(group, dim=1) if group > 1 else v
    qp = _positions(query_positions, batch, t, s - t, q.device)
    kp = _positions(key_positions, batch, s, 0, q.device)
    q32 = q.float() / math.sqrt(d)
    m = torch.full((batch, q_heads, t, 1), float("-inf"), device=q.device)
    l = torch.zeros((batch, q_heads, t, 1), device=q.device)
    acc = torch.zeros((batch, q_heads, t, v.shape[-1]), device=q.device)
    for start in range(0, s, tile):
        stop = min(start + tile, s)
        scores = q32 @ k[:, :, start:stop].float().transpose(-2, -1)
        visible = kp[:, None, None, start:stop] <= qp[:, None, :, None]
        scores = scores.masked_fill(~visible, float("-inf"))
        m_new = torch.maximum(m, scores.amax(-1, keepdim=True))
        # A row that has seen no visible key yet keeps m = -inf; use 0 as a safe reference.
        reference = torch.where(torch.isfinite(m_new), m_new, torch.zeros_like(m_new))
        alpha = torch.exp(m - reference)               # exp(-inf) = 0 for the first visible tile
        p = torch.exp(scores - reference)
        l = alpha * l + p.sum(-1, keepdim=True)
        acc = alpha * acc + p @ v[:, :, start:stop].float()
        m = m_new
    if bool((l == 0).any()):
        raise ValueError("A query row has no visible key")
    return (acc / l).to(q.dtype)

Two edge cases need care. A query row may see no visible key in an early tile (with causal masking, the first tiles of a long query block). Then $m = -\infty$, and $e^{-\infty - (-\infty)}$ is NaN. The code uses a safe reference of 0 for rows that haven’t seen any score yet. And if a row never sees a key at all, the result is undefined, so the function raises.

The same recurrence for one query and one head, in C++ and Rust:

inline std::vector<float> online_attention(const std::vector<float>& q, const std::vector<float>& k,
                                           const std::vector<float>& v, size_t S, size_t d, size_t tile) {
    float m = -INFINITY, l = 0, scale = 1.f / std::sqrt(float(d));
    std::vector<float> acc(d, 0.f);
    for (size_t start = 0; start < S; start += tile) {
        size_t end = std::min(S, start + tile);
        std::vector<float> s(end - start);
        float m_new = m;
        for (size_t j = start; j < end; ++j) {
            s[j - start] = std::inner_product(q.begin(), q.end(), k.begin() + j * d, 0.f) * scale;
            m_new = std::max(m_new, s[j - start]);
        }
        float alpha = std::exp(m - m_new);      // rescale what was accumulated under the old max
        l *= alpha;
        for (float& a : acc) a *= alpha;
        for (size_t j = start; j < end; ++j) {
            float p = std::exp(s[j - start] - m_new);
            l += p;
            for (size_t e = 0; e < d; ++e) acc[e] += p * v[j * d + e];
        }
        m = m_new;
    }
    for (float& a : acc) a /= l;
    return acc;
}
#![allow(unused)]
fn main() {
/// One query, one head, keys processed in tiles of `tile` with a running max m, running
/// denominator l and an unnormalized accumulator acc (Chapter 15). Equals softmax(qK^T)V.
pub fn online_attention(q: &[f32], k: &[f32], v: &[f32], s: usize, d: usize, tile: usize) -> Vec<f32> {
    let scale = 1.0 / (d as f32).sqrt();
    let (mut m, mut l) = (f32::NEG_INFINITY, 0.0f32);
    let mut acc = vec![0.0f32; d];
    for start in (0..s).step_by(tile) {
        let end = (start + tile).min(s);
        let scores: Vec<f32> = (start..end)
            .map(|j| q.iter().zip(&k[j * d..][..d]).map(|(a, b)| a * b).sum::<f32>() * scale)
            .collect();
        let m_new = scores.iter().cloned().fold(m, f32::max);
        let alpha = (m - m_new).exp(); // rescales everything accumulated so far
        l *= alpha;
        acc.iter_mut().for_each(|a| *a *= alpha);
        for (j, sc) in (start..end).zip(&scores) {
            let p = (sc - m_new).exp();
            l += p;
            for (a, vv) in acc.iter_mut().zip(&v[j * d..][..d]) {
                *a += p * vv;
            }
        }
        m = m_new;
    }
    acc.iter().map(|a| a / l).collect()
}
}

This PyTorch version saves memory but not time: each tile is still several separate kernels launched from a Python loop. The speed comes from doing a whole tile’s work inside one kernel, with the tile in on-chip memory.

The kernel

A FlashAttention forward kernel assigns each program a block of BLOCK_M queries for one (batch, head). The program loads its queries once, then streams BLOCK_N keys and values at a time through on-chip memory, applying the recurrence:

@triton.jit
def flash_fwd_kernel(q_ptr, k_ptr, v_ptr, o_ptr,
                     sqb, sqh, sqt, skb, skh, sks, svb, svh, svs, sob, soh, sot,
                     T, S, offset, heads, group, scale,
                     D: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
    """(Your engine: Chapter 15)"""
    pid_m = tl.program_id(0)                 # which tile of queries
    bh = tl.program_id(1)                    # which (batch, query head)
    b = bh // heads
    h = bh % heads
    kvh = h // group                         # GQA: several query heads read one KV head
    rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    rd = tl.arange(0, D)
    q = tl.load(q_ptr + b * sqb + h * sqh + rm[:, None] * sqt + rd[None, :], mask=rm[:, None] < T, other=0.0)
    q_pos = rm + offset                      # absolute position of each query row
    m = tl.full((BLOCK_M,), -float("inf"), tl.float32)
    l = tl.zeros((BLOCK_M,), tl.float32)
    acc = tl.zeros((BLOCK_M, D), tl.float32)
    # Causality: no key beyond the last query's position is ever needed.
    last_key = tl.minimum(S, offset + (pid_m + 1) * BLOCK_M)
    for start in range(0, last_key, BLOCK_N):
        rn = start + tl.arange(0, BLOCK_N)
        k = tl.load(k_ptr + b * skb + kvh * skh + rn[:, None] * sks + rd[None, :], mask=rn[:, None] < S, other=0.0)
        v = tl.load(v_ptr + b * svb + kvh * svh + rn[:, None] * svs + rd[None, :], mask=rn[:, None] < S, other=0.0)
        # input_precision="ieee": FP32 inputs get true FP32 math, not TF32 (BF16/FP16 are unaffected).
        s = tl.dot(q, tl.trans(k), input_precision="ieee").to(tl.float32) * scale   # [BLOCK_M, BLOCK_N]
        visible = (rn[None, :] <= q_pos[:, None]) & (rn[None, :] < S)
        s = tl.where(visible, s, -float("inf"))
        m_new = tl.maximum(m, tl.max(s, axis=1))
        m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)           # rows with nothing visible yet
        alpha = tl.exp(m - m_safe)                                        # rescale the old partial sums
        p = tl.exp(s - m_safe[:, None])
        l = alpha * l + tl.sum(p, axis=1)
        acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v, input_precision="ieee").to(tl.float32)
        m = m_new
    out = acc / l[:, None]
    tl.store(o_ptr + b * sob + h * soh + rm[:, None] * sot + rd[None, :], out.to(o_ptr.dtype.element_ty),
             mask=rm[:, None] < T)


def flash_attention(q, k, v, query_offset=None, block_m=32, block_n=32):
    """q [B, Hq, T, D], k/v [B, Hkv, S, D]; query i sits at absolute position offset + i
    (default offset = S - T, i.e. the queries are the newest tokens). D must be a power of two >= 16."""
    check_device(q, k, v)
    batch, heads, t, d = q.shape
    kv_heads, s = k.shape[1], k.shape[2]
    if heads % kv_heads or d & (d - 1) or d < 16:
        raise ValueError("Need Hq % Hkv == 0 and a power-of-two head dim >= 16")
    offset = s - t if query_offset is None else query_offset
    q, k, v = (x.contiguous() for x in (q, k, v))
    o = torch.empty_like(q)
    grid = (triton.cdiv(t, block_m), batch * heads)
    flash_fwd_kernel[grid](q, k, v, o, *q.stride()[:3], *k.stride()[:3], *v.stride()[:3], *o.stride()[:3],
                           t, s, offset, heads, heads // kv_heads, 1.0 / math.sqrt(d),
                           D=d, BLOCK_M=block_m, BLOCK_N=block_n)
    return o

Details worth noticing:

  • Causal block skipping. No key past the last query’s position is needed, so the loop stops at offset + (pid_m + 1) * BLOCK_M. That halves the work for causal attention.
  • Query offset. Query row i sits at absolute position offset + i. With offset = S - T, one kernel serves full prefill (T = S), chunked prefill (T < S) and decode (T = 1), exactly like your Chapter 5 position rule.
  • Grouped-query attention for free. Query head h reads KV head h // group. No repeated K/V tensors are materialized.
  • FP32 statistics, low-precision inputs. m, l and acc stay in FP32 registers. p is cast to the value dtype for the second tl.dot, so both matmuls run on tensor cores.
  • input_precision="ieee". On NVIDIA GPUs tl.dot runs FP32 inputs on TF32 tensor cores by default (Chapter 13’s 10-bit mantissa, about 1e-3 error), so an FP32 test against causal_attention fails at a 1e-4 tolerance even though the interpreter passes. "ieee" asks for true FP32 math. It changes only FP32 inputs: BF16 and FP16 still run on tensor cores.

How much memory traffic does this save? Each query block streams the visible K and V once, so the kernel reads about $2T^2 d/\text{BLOCK_M}$ elements of K and V. With $d = 128$ and BLOCK_M = 128, that’s about $2T^2$, against roughly $4T^2$ score elements moved by the unfused version (plus its own K/V reads). The bigger wins are elsewhere: nothing of size $T^2$ is ever allocated or written, and the time moves out of memory-bound elementwise kernels into two compute-bound matmuls per tile.

Note

Production FlashAttention adds much more. FlashAttention-2 reorders loops to parallelize over the sequence and cut non-matmul work. FlashAttention-3 uses Hopper’s asynchronous copies (TMA) and warp specialization (some warps load, others compute) and supports FP8. FlashInfer and vLLM’s kernels add paged KV caches (Chapter 25) and many variants. GPU Mode L12 (Thomas Viehmann’s FlashAttention lecture, with a from-scratch CUDA version) and L36 (FlashAttention-3 with CUTLASS) cover them.

Decode is different: flash-decoding

During decode, there’s one query per sequence and a long KV cache. One program per (batch, head) would leave most of the GPU idle: with batch 1 and 8 KV heads, that’s 8 programs for 132 SMs. Flash-decoding splits the KV sequence into chunks processed by different programs in parallel. Each produces a partial result $(m_i, \ell_i, a_i)$ for its chunk, and a second small kernel merges them with the same rescaling rule:

$$ m = \max_i m_i,\qquad \ell = \sum_i e^{m_i - m},\ell_i,\qquad a = \sum_i e^{m_i - m},a_i,\qquad O = a/\ell . $$

This merge rule is associative: any number of partial softmaxes can be combined in any grouping. That same property powers ring attention across multiple GPUs (GPU Mode L13) and the paged decode kernel in Chapter 25.

What FlashAttention does and doesn’t change

  • Changes: memory for intermediates goes from $O(T^2)$ to $O(T)$, and HBM traffic drops several-fold, so long-sequence attention becomes compute-bound and fast.
  • Doesn’t change: the result, up to rounding, and the arithmetic. Attention still costs $O(T^2 d)$ FLOPs. At very long contexts that quadratic arithmetic itself becomes the bottleneck, which is the motivation for Chapters 28 and 29.

Build it

Engine milestone 15: FlashAttention. Implement online_attention in engine/attention.py and flash_fwd_kernel in engine/kernels/triton_flash.py (the launcher flash_attention is provided).

pytest tests/test_ch15_flash.py

The tests try tile sizes from 1 to 64 (the answer must not depend on the tile size), rectangular queries with GQA, the worked example above, and the Triton kernel on prefill, chunked-prefill and decode shapes. On a GPU, benchmark your kernel against F.scaled_dot_product_attention for T = 512 to 8,192.

Stretch exercises

  1. ★ Instrument causal_attention and online_attention with torch.cuda.max_memory_allocated() for T = 4,096. How much memory does the online version save? Where: experiments/ch15.py (create it), importing both functions from engine.attention.
  2. ★★ Implement the flash-decoding merge: split online_attention’s key range into 4 chunks, compute each chunk’s $(m, \ell, a)$ separately, and merge. Verify against the unsplit result. Where: add a split/merge attention helper beside online_attention in engine/attention.py.
  3. ★★ Add a sliding-window option to the Triton kernel: skip tiles entirely before q_pos - window, and mask within the boundary tile. Where: flash_fwd_kernel and flash_attention in engine/kernels/triton_flash.py.
  4. ★★★ Write the backward pass of attention in PyTorch, using FlashAttention’s trick: recompute $P$ from the saved log-sum-exp per row instead of storing it. Check it against autograd. Where: add a backward helper in engine/attention.py; compare with autograd in tests/test_ch15_stretch.py.

Check your understanding

  1. Why is the accumulator $a$ kept unnormalized until the very end?
  2. Why is applying an ordinary softmax to each tile independently incorrect?
  3. Which complexity does FlashAttention reduce: arithmetic, intermediate memory, or both?
  4. Why does decode need a different parallelization (flash-decoding) than prefill?

Going deeper

  • PMPP §20.5 (pp. 492-503): FlashAttention, derived and implemented in CUDA; §20.6 (KV-cache arithmetic intensity) previews Chapter 16.
  • GPU Mode L12 (FlashAttention, with flash_attention.cu and a notebook), L13 (ring attention and the log-sum-exp merge, with howto_log_sum_exp.ipynb), L36 (FlashAttention-3).
  • Dao et al., FlashAttention (2022) and FlashAttention-2 (2023); Shah et al., FlashAttention-3 (2024); Milakov and Gimelshein, Online normalizer calculation for softmax (2018).
  • The Triton tutorial Fused Attention, which this chapter’s kernel simplifies.