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

32. Production attention kernels

In this chapter

  • Why the reference backend is slow, and what a fast one must do: one launch per layer for the whole mixed batch, K/V read straight from the blocks, each byte read once.
  • A unified Triton kernel for flattened batches: prefill chunks, decode tokens and verification blocks together, with grouped-query heads packed into one tile.
  • Split-KV ("flash-decoding") for long-context decode, and the log-sum-exp merge that makes it exact.
  • INT8 and FP8 KV caches with a scale per token and head: twice the context in the same memory.
  • When to use vendor kernels (FlashAttention, FlashInfer) and how to plug them in behind the same test.

You will build

unified_attention_kernel, splitkv_decode_kernel and merge_kernel in engine/kernels/triton_unified.py, and quantize_kv in engine/serve/triton_backend.py.

Time: 6-8 hours. GPU: recommended for timing (every test runs in Triton's interpreter on a CPU).

What the reference backend costs

Chapter 31’s ReferenceBackend is correct and slow, in three separate ways:

  1. It copies the context. For each request it gathers the request’s blocks into a contiguous tensor, then attention reads that copy. Every byte of K and V crosses the memory bus three times (read from the pool, written to the copy, read again) instead of once.
  2. It launches per request. A Python loop over requests issues a handful of kernels each. With 64 decoding requests that’s hundreds of launches per layer, and Chapter 19 showed what launch overhead does to decode.
  3. It reads K/V once per query head. causal_attention repeats each KV head for its group of query heads (repeat_interleave). Qwen3-8B has 4 query heads per KV head, so the cache is read four times.

How fast could it be? Decode attention is memory-bound: each step must read every cached key and value once, and does about 4 FLOPs per element read (two for $q \cdot k$, two for $p \cdot v$), far below an H100’s ~300 FLOPs per byte. So its time is bytes divided by bandwidth. Take 64 requests at 4,096 tokens of context on Qwen3-8B. One token of one layer holds $2 \times 8 \times 128 = 2{,}048$ values, 4 KiB in BF16, so one layer’s cache for the batch is $64 \times 4{,}096 \times 4$ KiB = 1.07 GB, and all 36 layers hold 38.7 GB. At 3.35 TB/s that’s 11.5 ms per decode step for attention alone, a floor no kernel can beat, against about 4.9 ms to read the 16.4 GB of weights. At this batch size and context length attention dominates the step, and a kernel that reads the cache more than once makes it proportionally slower.

A production attention kernel therefore has three goals: read each cached byte once, in one launch per layer for the whole batch, and keep enough programs running to use every SM.

One kernel for the whole batch

The flattened batch of Chapter 31 mixes requests in different states: a 512-token prefill chunk, thirty decode tokens, a 5-token verification block. A kernel written for one of these shapes is bad at the others. Prefill kernels tile over many query rows; decode kernels have one query row and must parallelize over something else. The unified kernel handles both by choosing its tiles from the metadata:

  • one program per (request, tile of BLOCK_Q query tokens, KV head);
  • each program loads the query rows of all GROUP query heads that share its KV head, so one tile of K and V serves all of them;
  • it walks the request’s block table one KV block per iteration, applying the causal rule by position and the online softmax of Chapter 15.
program (r, qt, g):   rows = BLOCK_Q tokens × GROUP_PAD heads
                      row m  ->  token qt*BLOCK_Q + m // GROUP_PAD,  head g*GROUP + m % GROUP_PAD

  q tile [BLOCK_Q·GROUP_PAD, D]      K block [BLOCK_SIZE, D]  (block = table[r, j])
        ┌──────────┐                       ┌──────┐
 tok 0  │ h0 h1 h2 h3 │  · Kᵀ  →  scores [rows, BLOCK_SIZE]  → online softmax → · V
 tok 1  │ h0 h1 h2 h3 │
  ...   └──────────┘

Packing the group into the tile’s rows is the GQA trick that FlashInfer and vLLM’s kernels use. It fixes problem 3 (K/V are loaded once per group), and it also makes decode tensor-core friendly: a single decode token has only one query row, but with 4 heads per group and BLOCK_Q = 4 the tile has 16 rows, the minimum tl.dot shape.

@triton.jit
def unified_attention_kernel(q_ptr, k_ptr, v_ptr, o_ptr, ks_ptr, vs_ptr,
                             qsl_ptr, seq_lens_ptr, table_ptr,
                             s_qt, s_qh, s_kb, s_ks, s_kh, s_sb, s_ss, s_table,
                             scale, window,
                             GROUP: tl.constexpr, GROUP_PAD: tl.constexpr, BLOCK_Q: tl.constexpr,
                             BLOCK_SIZE: tl.constexpr, D: tl.constexpr, QUANT: tl.constexpr):
    """(Your engine: Chapter 32)"""
    r, qt, g = tl.program_id(0), tl.program_id(1), tl.program_id(2)
    q_start = tl.load(qsl_ptr + r)
    q_len = tl.load(qsl_ptr + r + 1) - q_start
    length = tl.load(seq_lens_ptr + r)
    if qt * BLOCK_Q < q_len:
        # Row m of the tile is (query token m // GROUP_PAD, query head m % GROUP_PAD of group g).
        rows = tl.arange(0, BLOCK_Q * GROUP_PAD)
        token = qt * BLOCK_Q + rows // GROUP_PAD
        head = g * GROUP + rows % GROUP_PAD
        row_ok = (token < q_len) & (rows % GROUP_PAD < GROUP)
        d = tl.arange(0, D)
        q = tl.load(q_ptr + (q_start + token)[:, None] * s_qt + head[:, None] * s_qh + d[None, :],
                    mask=row_ok[:, None], other=0.0)
        if QUANT:
            q = q.to(tl.float32)
        q_pos = length - q_len + token                         # queries are the newest positions
        m_i = tl.full((BLOCK_Q * GROUP_PAD,), -float("inf"), tl.float32)
        l_i = tl.zeros((BLOCK_Q * GROUP_PAD,), tl.float32)
        acc = tl.zeros((BLOCK_Q * GROUP_PAD, D), tl.float32)
        # Keys needed: up to the tile's last query; with a sliding window, from its first query's window.
        last_key = tl.minimum(length, length - q_len + tl.minimum(q_len, (qt + 1) * BLOCK_Q))
        first_key = tl.where(window > 0, tl.maximum(length - q_len + qt * BLOCK_Q - window + 1, 0), 0)
        slots = tl.arange(0, BLOCK_SIZE)
        for j in range(first_key // BLOCK_SIZE, tl.cdiv(last_key, BLOCK_SIZE)):
            block = tl.load(table_ptr + r * s_table + j).to(tl.int64)
            base = block * s_kb + g * s_kh
            k = tl.load(k_ptr + base + slots[:, None] * s_ks + d[None, :])
            v = tl.load(v_ptr + base + slots[:, None] * s_ks + d[None, :])
            if QUANT:                                          # dequantize the tile in registers
                k = k.to(tl.float32) * tl.load(ks_ptr + block * s_sb + slots * s_ss + g)[:, None]
                v = v.to(tl.float32) * tl.load(vs_ptr + block * s_sb + slots * s_ss + g)[:, None]
            s = tl.dot(q, tl.trans(k), input_precision="ieee").to(tl.float32) * scale
            key = j * BLOCK_SIZE + slots
            visible = (key[None, :] <= q_pos[:, None]) & (key[None, :] < length)
            visible = visible & ((window <= 0) | (key[None, :] > q_pos[:, None] - window))
            s = tl.where(visible, s, -float("inf"))
            m_new = tl.maximum(m_i, tl.max(s, axis=1))
            m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)
            alpha = tl.exp(m_i - m_safe)
            p = tl.exp(s - m_safe[:, None])
            l_i = alpha * l_i + tl.sum(p, axis=1)
            acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v, input_precision="ieee").to(tl.float32)
            m_i = m_new
        out = acc / tl.where(l_i > 0, l_i, 1.0)[:, None]
        tl.store(o_ptr + (q_start + token)[:, None] * s_qt + head[:, None] * s_qh + d[None, :],
                 out.to(o_ptr.dtype.element_ty), mask=row_ok[:, None])

Points to notice:

  • Positions come from lengths, not from the input. Request $r$’s $t$ queries are the last $t$ of its seq_len positions, the same rule as the reference backend, so q_pos = length - q_len + token.
  • The key loop stops at the tile’s last query. For a prefill chunk, a tile of early queries never loads keys after its last row’s position. For a decode token, the loop covers exactly the context.
  • Programs with nothing to do exit at once. The grid’s second dimension is sized for the longest query in the batch (max_query_len), so a decode request’s programs for tiles 1, 2, … see qt * BLOCK_Q >= q_len and return. That wastes a few launches’ worth of scheduling, not memory traffic.
  • Sliding windows (Chapter 42’s Mistral and Gemma layers) need two changes: the loop starts at the first block inside the oldest query’s window, and the mask also drops keys older than window positions.
  • No contiguous copy. K and V tiles are loaded straight from the block with block * stride + slot * stride + head * stride; the block table is the only indirection.

The launcher computes the grid and the packing:

def unified_attention(q, k_cache, v_cache, meta, scale=None, window=0, k_scale=None, v_scale=None, block_q=None):
    """q [N, Hq, D] -> out [N, Hq, D] for a flattened batch (BatchMeta from serve/batch.py)."""
    check_device(q, k_cache, v_cache)
    n, heads, d = q.shape
    kv_heads, block_size = k_cache.shape[2], k_cache.shape[1]
    if heads % kv_heads or d & (d - 1) or block_size & (block_size - 1):
        raise ValueError("Need Hq % Hkv == 0 and power-of-two head_dim and block_size")
    group = heads // kv_heads
    pad = _group_pad(group)
    block_q = block_q or max(1, 16 // pad)            # at least 16 rows per tile: tensor-core friendly
    out = torch.empty_like(q)
    quant = k_scale is not None
    ks = k_scale if quant else q.new_empty(1)
    vs = v_scale if quant else q.new_empty(1)
    grid = (meta.num_reqs, triton.cdiv(meta.max_query_len, block_q), kv_heads)
    unified_attention_kernel[grid](
        q, k_cache, v_cache, out, ks, vs, meta.query_start_loc, meta.seq_lens, meta.block_table,
        q.stride(0), q.stride(1), k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
        ks.stride(0) if quant else 0, ks.stride(1) if quant else 0, meta.block_table.stride(0),
        scale or 1.0 / math.sqrt(d), window,
        GROUP=group, GROUP_PAD=pad, BLOCK_Q=block_q, BLOCK_SIZE=block_size, D=d, QUANT=quant)
    return out

GROUP_PAD rounds the group up to a power of two, because Triton’s tl.arange needs one. A model with 6 query heads per KV head pads to 8, wasting a quarter of each tile’s rows, a cost that vendor kernels avoid with specialized code paths. For grouped heads, block sizes and head dimensions that are powers of two, nothing is wasted.

Long contexts: split the keys

The unified kernel’s parallelism for a decode batch is requests × KV heads. Four requests with 32,768-token contexts on Qwen3-8B give $4 \times 8 = 32$ programs, on a GPU with 132 SMs. Three quarters of the GPU idles while each program walks 2,048 blocks one after another. This is the long-context, small-batch case: a single user summarizing a book.

Split-KV, published as Flash-Decoding (Dao et al., 2023), adds a third grid dimension: each request’s context is cut into $S$ splits, and each program attends over its split only. Each produces a partial output $o_s$, normalized within its split, and its log-sum-exp $\ell_s = m_s + \log \sum_{j \in s} e^{x_j - m_s}$. The exact result is a weighted average:

$$ o = \sum_s w_s, o_s, \qquad w_s = \frac{e^{\ell_s}}{\sum_{s’} e^{\ell_{s’}}} = e^{\ell_s - \ell}, \qquad \ell = \log \sum_s e^{\ell_s}. $$

That’s the online-softmax recurrence of Chapter 15 applied to whole splits instead of tiles: each partial result carries the normalizer it was computed with, so they can be combined in any order. The same merge appears in cascade attention (attend to a shared prefix once for many requests, then merge with each request’s own suffix) and in ring attention (Chapter 41), where the splits live on different GPUs.

@triton.jit
def splitkv_decode_kernel(q_ptr, k_ptr, v_ptr, ks_ptr, vs_ptr, seq_lens_ptr, table_ptr, po_ptr, pl_ptr,
                          s_qt, s_qh, s_kb, s_ks, s_kh, s_sb, s_ss, s_table,
                          s_por, s_pos, s_poh, s_plr, s_pls,
                          scale, blocks_per_split,
                          GROUP: tl.constexpr, GROUP_PAD: tl.constexpr, BLOCK_SIZE: tl.constexpr,
                          D: tl.constexpr, QUANT: tl.constexpr):
    """One (request, KV head, split) of a decode batch: partial output and log-sum-exp.  (Your engine: Chapter 32)"""
    r, g, split = tl.program_id(0), tl.program_id(1), tl.program_id(2)
    length = tl.load(seq_lens_ptr + r)
    rows = tl.arange(0, GROUP_PAD)
    head = g * GROUP + rows
    row_ok = rows < GROUP
    d = tl.arange(0, D)
    q = tl.load(q_ptr + r * s_qt + head[:, None] * s_qh + d[None, :], mask=row_ok[:, None], other=0.0).to(tl.float32)
    m_i = tl.full((GROUP_PAD,), -float("inf"), tl.float32)
    l_i = tl.zeros((GROUP_PAD,), tl.float32)
    acc = tl.zeros((GROUP_PAD, D), tl.float32)
    first = split * blocks_per_split
    last = tl.minimum(first + blocks_per_split, tl.cdiv(length, BLOCK_SIZE))
    slots = tl.arange(0, BLOCK_SIZE)
    for j in range(first, last):
        block = tl.load(table_ptr + r * s_table + j).to(tl.int64)
        base = block * s_kb + g * s_kh
        k = tl.load(k_ptr + base + slots[:, None] * s_ks + d[None, :]).to(tl.float32)
        v = tl.load(v_ptr + base + slots[:, None] * s_ks + d[None, :]).to(tl.float32)
        if QUANT:
            k = k * tl.load(ks_ptr + block * s_sb + slots * s_ss + g)[:, None]
            v = v * tl.load(vs_ptr + block * s_sb + slots * s_ss + g)[:, None]
        s = tl.dot(q, tl.trans(k), input_precision="ieee") * scale          # [GROUP_PAD, BLOCK_SIZE]
        s = tl.where((j * BLOCK_SIZE + slots)[None, :] < length, s, -float("inf"))
        m_new = tl.maximum(m_i, tl.max(s, axis=1))
        m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)
        alpha = tl.exp(m_i - m_safe)
        p = tl.exp(s - m_safe[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        acc = acc * alpha[:, None] + tl.dot(p, v, input_precision="ieee")
        m_i = m_new
    out = acc / tl.where(l_i > 0, l_i, 1.0)[:, None]
    lse = tl.where(l_i > 0, m_i + tl.log(tl.where(l_i > 0, l_i, 1.0)), -float("inf"))   # empty split: -inf
    tl.store(po_ptr + r * s_por + split * s_pos + head[:, None] * s_poh + d[None, :], out, mask=row_ok[:, None])
    tl.store(pl_ptr + r * s_plr + split * s_pls + head, lse, mask=row_ok)


@triton.jit
def merge_kernel(po_ptr, pl_ptr, o_ptr, splits, s_por, s_pos, s_poh, s_plr, s_pls, s_ot, s_oh,
                 SPLITS_PAD: tl.constexpr, D: tl.constexpr):
    """out = sum_s exp(lse_s - lse) * out_s, with lse = logsumexp_s(lse_s).  (Your engine: Chapter 32)"""
    r, h = tl.program_id(0), tl.program_id(1)
    s = tl.arange(0, SPLITS_PAD)
    d = tl.arange(0, D)
    lse = tl.load(pl_ptr + r * s_plr + s * s_pls + h, mask=s < splits, other=-float("inf"))
    top = tl.max(lse, axis=0)
    top = tl.where(top == -float("inf"), 0.0, top)
    w = tl.exp(lse - top)                                       # 0 for empty splits
    parts = tl.load(po_ptr + r * s_por + s[:, None] * s_pos + h * s_poh + d[None, :], mask=(s < splits)[:, None], other=0.0)
    total = tl.sum(w, axis=0)
    out = tl.sum(w[:, None] * parts, axis=0) / tl.where(total > 0, total, 1.0)
    tl.store(o_ptr + r * s_ot + h * s_oh + d, out.to(o_ptr.dtype.element_ty))

An empty split (a short request in a batch with long ones, or a padding row) writes $\ell_s = -\infty$, whose weight is $e^{-\infty} = 0$. The merge program guards the all-empty case, so padding rows produce zeros rather than NaNs, which matters when a CUDA graph (Chapter 33) runs padded batches.

How many splits? Enough programs to fill the GPU, but not so many that each split is a few blocks and the merge costs more than it saves:

def choose_splits(num_reqs, kv_heads, max_blocks, target_programs=264, min_blocks_per_split=4, max_splits=32):
    """Enough programs to fill the GPU (about two per SM on a 132-SM H100), without splits so
    small that the merge costs more than it saves."""
    want = triton.cdiv(target_programs, max(num_reqs * kv_heads, 1))
    return max(1, min(want, max_splits, triton.cdiv(max_blocks, min_blocks_per_split)))


def split_kv_decode(q, k_cache, v_cache, meta, num_splits=None, scale=None, k_scale=None, v_scale=None):
    """Decode-only batch (one query row per request): q [R, Hq, D] -> out [R, Hq, D]."""
    check_device(q, k_cache, v_cache)
    reqs, heads, d = q.shape
    kv_heads, block_size = k_cache.shape[2], k_cache.shape[1]
    group = heads // kv_heads
    max_blocks = triton.cdiv(max(meta.max_seq_len, 1), block_size)
    splits = num_splits or choose_splits(reqs, kv_heads, max_blocks)
    per_split = triton.cdiv(max_blocks, splits)
    parts = torch.empty((reqs, splits, heads, d), device=q.device, dtype=torch.float32)
    lse = torch.empty((reqs, splits, heads), device=q.device, dtype=torch.float32)
    quant = k_scale is not None
    ks = k_scale if quant else parts
    vs = v_scale if quant else parts
    splitkv_decode_kernel[(reqs, kv_heads, splits)](
        q, k_cache, v_cache, ks, vs, meta.seq_lens, meta.block_table, parts, lse,
        q.stride(0), q.stride(1), k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
        ks.stride(0) if quant else 0, ks.stride(1) if quant else 0, meta.block_table.stride(0),
        parts.stride(0), parts.stride(1), parts.stride(2), lse.stride(0), lse.stride(1),
        scale or 1.0 / math.sqrt(d), per_split,
        GROUP=group, GROUP_PAD=max(_group_pad(group), 16), BLOCK_SIZE=block_size, D=d, QUANT=quant)
    out = torch.empty_like(q)
    merge_kernel[(reqs, heads)](parts, lse, out, splits, parts.stride(0), parts.stride(1), parts.stride(2),
                                lse.stride(0), lse.stride(1), out.stride(0), out.stride(1),
                                SPLITS_PAD=max(_group_pad(splits), 2), D=d)
    return out

The backend uses split-KV only for decode-only batches with few (request, KV head) pairs and contexts of at least 2,048 tokens. For a batch of 64 decoding requests, the unified kernel already has $64 \times 8 = 512$ programs, and splitting would only add the merge.

Quantized KV caches

At long contexts the KV cache dominates both memory and decode time (Chapter 16), so storing it in 8 bits instead of 16 doubles the context that fits and nearly halves attention’s memory traffic. Two 8-bit formats are common:

formatvaluesstrengths
INT8 with a scale255 evenly spaced levels in $[-127s, 127s]$uniform precision; works on every GPU
FP8 E4M3 with a scale4 exponent bits, 3 mantissa bits, max 448relative precision across a wide range; native in Hopper and later tensor cores

The scale’s granularity matters more than the format. One scale per layer (vLLM’s FP8 KV cache, which uses 1.0 unless the checkpoint ships calibrated scales) is cheapest, but a single token or channel with a large value then wastes most of the code range for every other entry. This chapter uses one scale per (token, KV head): 4 extra bytes per 128-value head vector, a 3% overhead, which follows outliers that vary from token to token.

def quantize_kv(x, dtype):
    """x [N, Hkv, D] -> (codes [N, Hkv, D], scales [N, Hkv]) with one scale per token and head.  (Your engine: Chapter 32)

    Symmetric: scale = max|x| / q_max, so the largest element maps to the code range's edge
    (127 for int8, 448 for fp8 E4M3). Per-token scales follow outliers that differ token
    to token; per-head scales keep one head's large channels from crushing another head.
    """
    code_dtype, q_max = KV_DTYPES[dtype]
    amax = x.float().abs().amax(-1).clamp_min(1e-8)
    scale = amax / q_max
    scaled = x.float() / scale[..., None]
    codes = scaled.round().clamp(-127, 127).to(code_dtype) if dtype == "int8" else scaled.to(code_dtype)
    return codes, scale

Quantization happens on write, once per token. Dequantization happens in the kernel’s registers: the tile of codes is loaded (half the bytes), converted and multiplied by its per-slot scale, then used exactly like a BF16 tile. Both kernels take a QUANT constant that switches this on, so the same code serves both pools.

class TritonBackend(ReferenceBackend):
    name = "triton"

    def __init__(self, kv_cache_dtype="auto", split_kv="auto"):
        if kv_cache_dtype not in ("auto", *KV_DTYPES):
            raise ValueError(f"kv_cache_dtype must be auto, int8 or fp8, not {kv_cache_dtype!r}")
        self.kv_cache_dtype, self.split_kv = kv_cache_dtype, split_kv

    @property
    def quantized(self):
        return self.kv_cache_dtype != "auto"

    def allocate(self, shape, dtype, device):
        if not self.quantized:
            return super().allocate(shape, dtype, device)
        code_dtype = KV_DTYPES[self.kv_cache_dtype][0]
        k, v = (torch.zeros(shape, dtype=code_dtype, device=device) for _ in range(2))
        k_scale, v_scale = (torch.zeros(shape[:3], dtype=torch.float32, device=device) for _ in range(2))
        return k, v, k_scale, v_scale

    def write(self, cache, k, v, meta):
        if not self.quantized:
            return write_kv(cache[0], cache[1], k, v, meta.slot_mapping)
        (kc, ks), (vc, vs) = quantize_kv(k, self.kv_cache_dtype), quantize_kv(v, self.kv_cache_dtype)
        bits = (lambda t: t.view(torch.uint8)) if self.kv_cache_dtype == "fp8" else (lambda t: t)
        write_kv(bits(cache[0]), bits(cache[1]), bits(kc), bits(vc), meta.slot_mapping)   # same bytes, any dtype
        write_kv(cache[2][..., None], cache[3][..., None], ks[..., None], vs[..., None], meta.slot_mapping)

    def use_split_kv(self, meta, kv_heads):
        """Split long decode contexts when one program per (request, KV head) can't fill the GPU."""
        if self.split_kv != "auto":
            return bool(self.split_kv) and meta.decode_only
        return meta.decode_only and meta.num_reqs * kv_heads < 128 and meta.max_seq_len >= 2048

    def forward(self, q, cache, meta, scale=None, window=0):
        from ..kernels.triton_unified import split_kv_decode, unified_attention
        scales = (cache[2], cache[3]) if self.quantized else (None, None)
        if window == 0 and self.use_split_kv(meta, cache[0].shape[2]):
            return split_kv_decode(q, cache[0], cache[1], meta, None, scale, *scales)
        return unified_attention(q, cache[0], cache[1], meta, scale, window, *scales)

FP8 has a catch on the CPU: PyTorch’s CPU kernels don’t implement index_copy_ for float8_e4m3fn. The write therefore copies the same bytes through a uint8 view. A dtype is just an interpretation of bytes; for a copy, the interpretation doesn’t matter.

How much accuracy does it cost? Keys are more sensitive than values, because errors in $q \cdot k$ pass through the exponential of the softmax, and real models’ keys have a few large outlier channels. KIVI (Liu et al., 2024) quantizes keys per channel and values per token for this reason; KVQuant (Hooper et al., 2024) goes further with non-uniform codes. Per-token-and-head scales at 8 bits are a robust default; measure with the model’s own perplexity before going lower (stretch exercise 3).

Vendor kernels

Should a production engine write its own attention kernels at all? Every major engine uses vendor kernels on its main path: vLLM uses FlashAttention 2/3 and FlashInfer on NVIDIA GPUs, SGLang uses FlashInfer and its own Triton kernels, TensorRT-LLM uses NVIDIA’s hand-written kernels. On Hopper, FlashAttention-3 (Shah et al., 2024) reaches 75% of peak by using features that Triton exposes only partly: asynchronous TMA loads, warp-specialized producer and consumer warps, wgmma instructions, and FP8 compute with incoherent processing. For prefill on an H100, expect FA3 to be well ahead of a straightforward Triton kernel; for decode, which is memory-bound, a good Triton kernel gets much closer, because the limit is bandwidth rather than instruction scheduling.

Your kernels remain useful in three ways: they run where vendor kernels don’t (AMD GPUs through Triton, Chapter 39; new architectures before the vendor library supports them; the CPU interpreter for testing), they’re the reference you understand completely, and they’re the starting point for operations no library has yet (the stretch exercises). The engine’s backend interface makes vendor kernels a drop-in:

class FlashAttnBackend(ReferenceBackend):
    """Delegates to flash_attn_varlen_func(..., block_table=...). FlashAttention-2 requires the
    page (block) size to be a multiple of 256; vLLM's fork of it accepts 16."""
    name = "flash_attn"

    def __init__(self):
        from flash_attn import flash_attn_varlen_func          # noqa: F401  (fail early if missing)
        self.fn = flash_attn_varlen_func

    def forward(self, q, cache, meta, scale=None, window=0):
        k_cache, v_cache = cache[0], cache[1]
        if k_cache.shape[1] % 256:
            raise ValueError("flash-attn 2 paged attention needs block_size % 256 == 0")
        cu_k = torch.zeros(meta.num_reqs + 1, dtype=torch.int32, device=q.device)
        cu_k[1:] = torch.cumsum(meta.seq_lens, 0)
        return self.fn(q, k_cache, v_cache, meta.query_start_loc, cu_k, meta.max_query_len, meta.max_seq_len,
                       softmax_scale=scale or 1.0 / math.sqrt(q.shape[-1]), causal=True,
                       window_size=(window - 1, 0) if window else (-1, -1), block_table=meta.block_table)

The adapter passes the same metadata the Triton kernel uses. FlashAttention-2’s paged interface requires a page size divisible by 256, so you’d run the engine with block_size=256, which lowers prefix-cache granularity (vLLM ships a fork of FlashAttention that accepts 16). FlashInfer’s interface is different again: a plan call on the CPU partitions the batch’s work once per step, then run launches the kernel for each layer, which moves the load-balancing decisions out of the kernel. Whatever you plug in, it must pass the same test as your own kernel: equal to the reference backend on a mixed batch.

The tests and the measurements

The milestone tests compare the unified kernel with the reference backend on a batch with requests in every state (mid-prefill chunk, a full 16-token prompt, a 21-token prompt, decodes at 70 and 130 tokens), for grouped heads with group sizes 4 and 3 (padded to 4) and for plain multi-head attention; the sliding window; split-KV with 1, 3 and 8 splits; padding rows producing exact zeros; INT8 and FP8 caches matching the reference on the dequantized values and staying close to full precision; and the engine core producing every request’s solo output with the Triton backend.

python run.py backends

On a CPU, with Qwen3-8B’s attention shape (32 query heads, 8 KV heads, head dim 128) and contexts shrunk 8× so that the interpreter finishes in a few minutes, it prints the errors against the reference:

{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton unified", "max_abs_error": 7.152557373046875e-07, "ms": null}
{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton, int8 KV", "max_abs_error": 0.009663641452789307, "ms": null}
{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton, fp8 KV", "max_abs_error": 0.03983457386493683, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton unified", "max_abs_error": 1.564621925354004e-07, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton split-kv", "max_abs_error": 1.4156103134155273e-07, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton, int8 KV", "max_abs_error": 0.0009696260094642639, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton, fp8 KV", "max_abs_error": 0.005297387018799782, "ms": null}
{"kv_cache": "bf16", "qwen3_8b_bytes_per_token": 147456, "tokens_per_GB": 6781}
{"kv_cache": "int8", "qwen3_8b_bytes_per_token": 76032, "tokens_per_GB": 13152}

The exact kernels agree with the reference to FP32 rounding. INT8 is four to five times more accurate than FP8 at the same size, because E4M3 has only 3 mantissa bits, and both errors shrink with longer contexts as the softmax averages over more values. On a GPU, the ms column times each backend; compare each against the bandwidth floor you computed at the start of the chapter (bytes of K and V read, divided by your GPU’s measured bandwidth from Chapter 10).

Build it

Engine milestone 32: production attention. Implement unified_attention_kernel, splitkv_decode_kernel and merge_kernel in engine/kernels/triton_unified.py, and quantize_kv in engine/serve/triton_backend.py (the launchers, the split heuristic, the backend class and the FlashAttention adapter are provided).

pytest tests/test_ch32_attention_kernels.py
python run.py backends --impl engine

Then serve with it: EngineConfig(attention_backend="triton", block_size=16), and kv_cache_dtype="int8" or "fp8" for a quantized cache.

Stretch exercises

  1. ★ Autotune the unified kernel with @triton.autotune over BLOCK_Q, num_warps and num_stages, keyed on max_query_len and the group size. Plot achieved bandwidth for decode batches of 1-256 requests at 4,096 tokens against the floor. Where: unified_attention_kernel and its launcher in engine/kernels/triton_unified.py.
  2. ★★ Implement cascade attention: for a batch whose requests share a long cached prefix, attend all of their queries to the shared blocks in one tile-friendly pass (the prefix’s K/V read once for the whole batch), attend each to its own suffix, and merge with merge_kernel. Measure the saving with 64 requests sharing a 4,000-token system prompt. Where: new cascade kernels in engine/kernels/triton_unified.py, dispatched by TritonBackend.forward in engine/serve/triton_backend.py.
  3. ★★ Measure KV quantization’s effect on quality: perplexity of Qwen3-0.6B on a held-out text with BF16, INT8 and FP8 caches, then per-channel keys (KIVI) at 4 bits. Where: experiments/ch32.py (create it) for BF16/INT8/FP8; extend cache layout/read/write in engine/serve/triton_backend.py and engine/kernels/triton_unified.py for KIVI.
  4. ★★★ Add attention sinks (gpt-oss): one learned logit per head that joins every softmax’s denominator but contributes no value. It’s two lines in the online softmax: start m_i and l_i from the sink instead of $-\infty$ and 0. Test against a dense implementation. Where: online-softmax initialization in engine/kernels/triton_unified.py; add the dense comparison in engine/serve/attention.py.

Check your understanding

  1. Why is decode attention memory-bound, and what’s the minimum time for a step, given the context lengths and the bandwidth?
  2. What does packing the GQA group into a tile’s rows save, and why does it also help decode use tensor cores?
  3. Why does a decode-only batch with 4 long requests need split-KV, while one with 64 requests doesn’t?
  4. Derive the merge weights $w_s$ from the definition of softmax. Why can the splits be merged in any order?
  5. Why must an empty split write $-\infty$ as its log-sum-exp, and what would a padding row produce if the merge didn’t guard the all-empty case?
  6. Why are per-token-and-head scales more accurate than a single scale for the whole KV cache?

Going deeper

  • Dao, Haziza, Massa, Sizov, Flash-Decoding for long-context inference (Stanford CRFM, 2023); Shah et al., FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision (2024).
  • Ye et al., FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving (MLSys 2025), for load-balanced scheduling of ragged batches and cascade attention; GPU Mode L40 (FlashInfer).
  • vLLM’s Triton attention backend (vllm/attention/ops/ and vllm/v1/attention/backends/), which follows the unified design of this chapter and is one of the backends vLLM runs on AMD GPUs.
  • Liu et al., KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache (2024); Hooper et al., KVQuant (2024).
  • PMPP §20.5 (FlashAttention, pp. 492-503) and GPU Mode L12 (Flash Attention), for the tiling ideas this kernel builds on.