Many production batches are mostly the same text. A support assistant runs every request behind one 3,000-token system prompt. Best-of-N sampling asks for 16 answers to one prompt. An RL rollout step samples a group of 64 completions per question. Beam search and tree-of-thought fork one context into many. In each case dozens of sequences decode together on top of an identical prefix.

Prefix caching already solves half of this: the prefix is prefilled once and its KV blocks are stored once. This article is about the other half. Storing the prefix once does not mean reading it once. A standard decode kernel walks each sequence's block table on its own, so a prefix shared by B sequences is streamed from HBM B times on every step. Shared-prefix attention kernels (Hydragen, FlashInfer's cascade attention and the cascade path in vLLM's V1 engine) split attention at the prefix boundary so the prefix is read once per batch. You will see why that is exact, what it saves, a worked example, the code, and when it does not pay.

Storing once is not reading once

Keep two kinds of sharing apart. Storage sharing means one physical copy of the prefix KV, referenced by many sequences. It is done with hashed, reference-counted blocks and copy-on-write, covered in the paged KV cache article and, for the hit path and block lifetime, in prefix caching in depth. It saves memory and prefill compute.

Bandwidth sharing means the bytes of that copy cross the HBM bus once per step instead of once per sequence. Paged attention does not give you that. Each sequence's query is a single token, so its attention is a matrix-vector product: load every key and value of its context, do one multiply-add per loaded element per query head, move on. Arithmetic intensity is a few FLOPs per byte, so decode attention is purely memory-bound, and the time is the bytes. Some re-reads hit L2, but an H100's 50 MB L2 cannot hold a 512 MiB prefix, so most of them go to HBM.

The arithmetic of a decode step

Take a model in the Llama 3 8B shape: 32 layers, 8 KV heads, head dimension 128, fp16 KV. One token of KV costs 2 (K and V) x 32 x 8 x 128 x 2 bytes = 131,072 bytes, 128 KiB. A 4,096-token prefix is therefore 512 MiB of KV.

Now decode a batch of B sequences that share that prefix. Per step, per-sequence attention reads 512 MiB x B of prefix KV. Cascade attention reads 512 MiB once. Weights are read once per step regardless of B, which is about 16 GB for 8B parameters in bf16.

Batch BPrefix KV read, per-sequencePrefix KV read, cascadeTime at 3.35 TB/s (per-seq vs cascade)
84 GiB0.5 GiB1.3 ms vs 0.16 ms
3216 GiB0.5 GiB5.1 ms vs 0.16 ms
6432 GiB0.5 GiB10.3 ms vs 0.16 ms
12864 GiB0.5 GiB20.5 ms vs 0.16 ms

For comparison, reading the weights costs about 4.8 ms per step. At B = 64 the redundant prefix traffic is twice the weight traffic, so the step is dominated by re-reading a cache that never changes. That is the regime where Hydragen's authors report up to 32x end-to-end throughput on CodeLlama-13b and FlashInfer's cascade post reports up to 31x over vLLM's PagedAttention kernel for 32,768-token prompts at batch 256. Those are best cases with long prefixes and huge batches. Use the table to estimate your own case. Do not quote the headline.

Prefix KV bytes read per decode step (8B-class model, P = 4096, fp16 KV)B = 84 GiB read per step (per-sequence)0.5 GiB (cascade, read once)B = 3216 GiB read per step (per-sequence)0.5 GiB (cascade, read once)B = 6432 GiB read per step (per-sequence)0.5 GiB (cascade, read once)B = 12864 GiB read per step (per-sequence)0.5 GiB (cascade, read once)At about 3.35 TB/s of HBM bandwidth, 32 GiB is about 10 ms per step before any weight traffic.
Bars are the arithmetic above. Real kernels also read the suffix KV and the weights, which do not change between the two schemes.

Splitting attention at the prefix boundary

Why can you split attention at all? Softmax over a concatenated key set decomposes. For one query q and keys split into a prefix part A and a suffix part B, compute each part's output OA, OB (each normalised over its own keys) and its log-sum-exp LSEA, LSEB of the scaled scores. The full output is a weighted average, O = (eLSE_A OA + eLSE_B OB) / (eLSE_A + eLSE_B), computed with the maximum subtracted for stability. This is the same rescaling FlashAttention does between tiles, applied between two kernels. The result is exact up to floating-point rounding.

The win comes from pass 1. All B queries attend to the same prefix keys, so stack them into a [B x d] matrix and the prefix pass becomes one matrix-matrix product against [d x P]. That is tensor-core work with high arithmetic intensity, and the prefix KV is loaded once. With grouped-query attention the rows multiply further: each KV head serves its group of query heads, so the effective M dimension is B times the group size. Pass 2 is ordinary paged decode over each sequence's short private suffix. A small kernel merges the two (output, LSE) pairs.

This NumPy check reproduces the decomposition and compares it with single-pass attention:

import numpy as np
rng = np.random.default_rng(0)
d, P, B, S = 64, 512, 8, 40                 # head dim, prefix, batch, suffix len
Kp, Vp = rng.standard_normal((P, d)), rng.standard_normal((P, d))
Ks, Vs = rng.standard_normal((B, S, d)), rng.standard_normal((B, S, d))
Q = rng.standard_normal((B, d))
scale = 1 / np.sqrt(d)

def attn(q, K, V):                          # returns output and log-sum-exp
    s = q @ K.T * scale
    m = s.max(-1, keepdims=True); e = np.exp(s - m); l = e.sum(-1, keepdims=True)
    return (e @ V) / l, (m + np.log(l))[..., 0]

def merge(o1, l1, o2, l2):
    m = np.maximum(l1, l2); w1, w2 = np.exp(l1 - m), np.exp(l2 - m)
    return (w1[:, None] * o1 + w2[:, None] * o2) / (w1 + w2)[:, None]

o1, l1 = attn(Q, Kp, Vp)                    # pass 1: ONE matmul for the whole batch
o2 = np.zeros((B, d)); l2 = np.zeros(B)
for b in range(B):                          # pass 2: private suffixes
    o, l = attn(Q[b:b+1], Ks[b], Vs[b]); o2[b], l2[b] = o[0], l[0]

ref = np.stack([attn(Q[b:b+1], np.vstack([Kp, Ks[b]]), np.vstack([Vp, Vs[b]]))[0][0]
                for b in range(B)])
print(np.abs(merge(o1, l1, o2, l2) - ref).max())   # 3.05e-16 in float64
Cascade attention: the shared prefix is read once per batch, suffixes once per sequenceBatch of B queriesone new token per sequencePrefix KV (one copy)P tokens, sharedSuffix KV x Bprivate tokens per sequencePass 1: prefix[B x d] @ [d x P]: one GEMMPass 2: suffixB small attentionsO1, LSE1O2, LSE2Merge by log-sum-expexact, not approximateAttention outputidentical to one passPaged attention with prefix caching stores the prefix once but still reads it B times per step.Cascade attention changes the read count, which is what decode time is made of.
The two passes and the merge. Causal masking is trivial here because every suffix position sees the whole prefix.

Where it is implemented

Three implementations are worth knowing:

  • Hydragen (Juravsky et al., ICML 2024) is the reference formulation: attention split into prefix and suffix, inter-sequence batching of the prefix queries, and the LSE merge. It also extends to tree-shaped sharing.
  • FlashInfer cascade attention. flashinfer.cascade.merge_state merges two (output, LSE) pairs, merge_states merges many, and MultiLevelCascadeAttentionWrapper runs multi-level cascades over a unified paged KV table. Its documentation warns that more levels are not always better because every merge has a cost.
  • vLLM V1 has a cascade attention path. A heuristic chooses it when the batch has a long enough common prefix, and --disable-cascade-attn turns it off. The thresholds are internal and change between releases. Check the source of the version you run before relying on a specific cut-over.

Note what all three need: the shared prefix must be the same physical blocks. Cascade attention works on top of storage sharing. If prefix caching missed, or two requests carry byte-different system prompts, there is no common prefix to cascade over.

Prefix trees: multi-level sharing

Real sharing is a tree, not one prefix. Picture a system prompt shared by everyone, a 20-page document shared by one user's questions, and then 16 samples per question. A two-level cascade attends the system prompt across the whole batch, the document across that user's sequences, and then the private suffixes, with a merge after each level. Hydragen's tree variant cut inference time by 55% on competitive-programming problems where many samples fork from one problem statement.

The common workloads that produce this shape are parallel sampling (n > 1), best-of-N reranking, beam search, self-consistency voting, and RL rollouts such as GRPO, which deliberately sample a group of completions per prompt. Rollout generation is often the most expensive phase of RL post-training, so it is the place most likely to repay the work.

Worked example: 64 rollouts per prompt

Worked example: 64 rollouts per prompt. An 8B policy generates 64 completions for each of 2 prompts. The prompts are 4,096 tokens and each completion runs to 1,024 tokens. The batch is 128 sequences in two groups of 64.

  1. Per-sequence attention, by the final step: each sequence reads 4,096 + 1,024 tokens of KV, 128 x 5,120 x 128 KiB = 80 GiB per step, about 25.6 ms at 3.35 TB/s, plus 4.8 ms of weights.
  2. Cascade: prefix reads drop to 2 x 512 MiB = 1 GiB. Suffix reads stay at 128 x 1,024 x 128 KiB = 16 GiB. The total is 17 GiB, about 5.4 ms, plus 4.8 ms of weights.
  3. The step time at the end of generation goes from about 30 ms to about 10 ms. Early steps, where suffixes are short, improve even more. The average gain over the run is roughly 3-4x on decode, which shrinks the generation share of each RL step.

Notice the limit. Once suffixes get long relative to the prefix, private KV dominates again and cascade cannot help. Shared-prefix attention pays most for long shared contexts with short answers, and least for short prompts with long reasoning traces.

Getting shared batches from the scheduler

The kernel can only exploit sharing that the scheduler puts in the same batch:

  • Submit forks together. Use one request with n=64 rather than 64 independent requests, so the engine knows they are siblings and keeps them on one replica.
  • Route by prefix. Across replicas, send requests with the same prefix to the replica that holds it. The KV-aware routing in LLM routing strategies is the precondition for any cross-request sharing.
  • Keep prefixes byte-identical. Put timestamps, user names and request IDs after the shared part, never inside the system prompt.
  • Watch preemption. If the engine evicts sibling sequences under memory pressure, the batch loses its common prefix and silently falls back to per-sequence attention. Size the pool with KV cache sizing so a full group fits.

Failure modes

Failure modeSymptomFix
Short prefix or small batchCascade slower than plain decodeLeave the engine heuristic on; benchmark below the cut-over
Prefixes differ by one byteNo speed-up despite identical-looking promptsHash the rendered prompt; move volatile fields to the end
Mixed batchCommon prefix of the whole batch is near zeroUse multi-level cascade or group siblings into one batch
LSE kept in low precisionSmall output drift against the referenceKeep LSE and the merge in fp32; diff against single-pass attention
Too many cascade levelsMerge kernels dominate the stepCollapse levels that save under a few MiB per step
Sliding-window or chunked-local layersPrefix not fully visible to every queryApply cascade only on global-attention layers or disable it

Trade-offs

Cascade attention adds a second kernel launch and a merge for every attention layer, plus scheduling logic to find the common prefix. With a prefix under a few hundred tokens, or only a handful of sequences, that overhead can cost more than it saves. With thousands of tokens and dozens of sequences it removes most of the attention time. It also changes the floating-point summation order, so outputs are not bit-identical to single-pass attention. Greedy decoding can occasionally take a different token at a near-tie. If you need bit-for-bit reproducibility, for example in an evaluation harness, disable it. Otherwise it is the same mathematical function.

What to do next

  1. Measure your sharing: log the common-prefix length and sibling count per decode batch for one day.
  2. Run the arithmetic: per-token KV bytes x prefix x batch, against your weight bytes. If the prefix term is not comparable to the weights, stop here.
  3. Confirm prefix caching hits first; cascade needs shared physical blocks.
  4. Submit forks as one n > 1 request and route same-prefix traffic to one replica.
  5. A/B the engine with cascade on and with --disable-cascade-attn and compare decode step time, tokens per second and output diffs on a fixed prompt set.
  6. For RL rollouts, measure generation time per step before and after. That is the number that moves.
Key takeaway: Prefix caching stores a shared prefix once, but per-sequence decode kernels still read it once per sequence on every step. Cascade attention splits softmax at the prefix boundary. It runs the prefix as one batched matrix product, the suffixes separately, and merges the two exactly by log-sum-exp. The saving is prefix bytes x (batch - 1) per step. It is decisive for long shared contexts with many siblings, such as parallel sampling and RL rollouts, and negligible for short prompts.