三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

大模型推理加速:共享前缀(Shared Prefix)优化技术的全面解析 - -银光

大模型推理加速:共享前缀(Shared Prefix)优化技术的全面解析 - -银光

本文已于 2026.08.15 发表于公众号和知乎。

1. 核心原理:为什么前缀可以共享

1.1 研究动机

提到共享前缀优化,大家很容易想到 vLLM/SGLang 等引擎的 prefix KV cache 机制,但这只是跨批次复用前缀的场景。如果在同一个批次内多个请求共享前缀,是否也有针对性的优化手段?本文将系统回答以下三个问题:

  • prefill 阶段:同批次多个请求共享前缀,除了拆分到不同批次命中 prefix cache 之外,还有哪些方案可以减少重复计算?
  • decode 阶段:KV cache 已存在、不存在重复计算,但多个请求仍会从 HBM 重复读取相同的共享前缀 KV——这种带宽浪费如何优化?
  • 反常识:减少了重复计算或 HBM 读取,性能就一定更好吗?

1.2 为什么前缀可以共享

Decoder-only 架构的注意力计算采用因果掩码(causal mask):序列中每个 token 只能看到其左侧的前序 token。这一特性使 KV cache 可以前缀共享——每个 token 的 KV 只由其前序 token 决定,因此前缀相同的 prompt,其前缀部分的 KV 必然一致,可被多个请求复用。

2. 三大应用场景分类

2.1 三种共享前缀场景

从系统角度,共享前缀优化涵盖三类场景:

场景阶段条件本质
场景 1 Prefill 不同批次 当前批次复用历史批次已缓存的 KV,避免重复 prefill
场景 2 Prefill 相同批次 多请求共享同一前缀,前缀 KV 尚未生成,需要减少 prefill 阶段的重复计算
场景 3 Decode 相同批次 多请求共享同一前缀,前缀 KV 已存在,需要减少 decode 阶段的重复 HBM 读取

场景 1 是成熟技术,对应第 3 章;场景 2 对应第 4 章,场景 3 对应第 6 章。

3. 场景 1:跨批次前缀共享

跨批次前缀共享(场景 1)是成熟且广泛应用的技术,尤其适合多轮对话,对降低推理成本效果显著。观察各家大模型 API 公司的定价即可感知:命中前缀 cache 的 token 单价远低于未命中的。譬如百万输入 token:

模型命中缓存未命中缓存
DeepSeek V4 Pro ¥0.022 ¥0.66
GLM-5.2 $0.26 $1.4
Kimi K3 $0.30 $3.00

命中缓存的部分序列几乎不含计算成本,只有 KVCache 的存储和读取(传输)的成本。因此多轮对话(尤其 agent 场景)中,通常追加聊天信息而非修改前几轮对话,使历史前缀持续命中,并随对话轮数累积不断变长。

跨批次前缀共享涉及庞大的技术体系,主要包括:

  • 前缀匹配机制:基于 radix tree(SGLang)或块级 token rolling hash(vLLM)的 cache 存储与查找。
  • 多级存储:HBM → DRAM → NVMe → 远程存储,逐级容量递增、带宽递减、延迟递增。代表系统有 LMCache、Mooncake 等,随着长上下文应用普及,显得愈发重要。
  • Fetch vs Recompute 权衡:从慢速介质取回 KV 的 IO 成本 vs 重新 prefill 的计算成本,结合 cache 感知调度做决策。
  • 计算与 IO 重叠:为缓解 IO 延迟引入 KV cache 预读、KV cache 分层 overlap、offload 回写、传输压缩等技术。
  • KV cache 压缩/量化:FP8/INT4 量化、低秩压缩(如 MLA)。

上述每一项技术都值得单独专题展开,社区也有丰富的介绍文章,本文仅做概览。第 4、5 章聚焦于同批次内的共享前缀优化(场景 2 和场景 3)。

4. 场景 2:同批次 Prefill 的五种技术路线

当同一 batch 中多个请求拥有相同前缀时,经我的系统分析调查,可分为 5 种处理方案:

方案思路实现
方案 1 朴素实现(无优化) 每个请求独立完整 prefill
方案 2 批次拆分 + Cache 首个请求完整 prefill 生成前缀 KV,后续请求只 prefill 各自后缀并复用前缀 KV(跨批次拆分)
方案 3 Direct Ragged Prefill 打平拼接为一次 batch prefill kernel,共同前缀 Q 只算一次
方案 4 Cascade Attention 先算前缀(生成前缀 KV),再分两级算后缀(shared/unique),最后 merge
方案 5 掩码合并(kMultiItemScoring / custom_mask) 将多请求合并为一个逻辑请求,用 mask 隔离各 item 的可见范围(kMultiItemScoring:结构化 mask;custom_mask:任意形状 mask)

注:Direct Ragged Prefill:利用 ragged layout 做同批次前缀去重,具体是将多个变长请求的 Q 打平拼接为一次 batch prefill kernel 调用,通过 qo_indptr 标记各请求边界,结合 kv indices 标记,使共同前缀的 Q 只计算一次。

下面分别介绍这五种方案。先介绍各方案在"请求 / 序列"层面的浅层原理,暂不涉及算子层细节;等五种方案全部介绍完之后,再在 5.1 节深入算子层,从 CTA(Cooperative Thread Array,即 thread block)视角重新审视各方案。

4.1 方案 1:朴素实现(无优化)

每个请求独立完整 prefill,不做任何共享前缀优化。它也是后续各优化方案的 baseline。

底层原理

以两个 100-token 请求为例,共享 96-token 前缀,后缀各 4-token:

P  = 96-token 共同前缀
S0 = req0 的 4-token 独有后缀
S1 = req1 的 4-token 独有后缀

req0 = [P, S0],长度 100
req1 = [P, S1],长度 100

朴素实现下,两个请求互不感知,各自作为独立请求完成一次完整的 causal prefill:

request 0: Q = [P, S0],KV = [P, S0]   → 100 个 query 对 100 个 KV 做 causal 注意力
request 1: Q = [P, S1],KV = [P, S1] → 100 个 query 对 100 个 KV 做 causal 注意力

计算示意图

4.1

如上图所示,两个请求各自独立完成 prefill,共同前缀 P 的计算结果与生成的 KV cache 都是重复的。

4.2 方案 2:SGLang 的批次拆分

转为不同批次的请求,然后通过缓存消除共同前缀的 prefill 计算,下面以 SGLang 实现的 in-batch prefix caching 为例展开介绍。

底层原理

SGLang 用 radix tree 记录已完成序列的 KVCache,对新到来的请求,在执行推理之前,对每个请求做最长前缀匹配,匹配上的序列复用已有 KVCache,这是跨批次的前缀共享。对于同一个 batch 里的不同请求有共同前缀的情况,SGLang 提供 lpm(longest prefix match)调度策略,将请求拆分为两批次,第一个请求完整 prefill(把公共前缀的 KV 写入 cache),后续请求命中 cache、只 prefill 后缀。还是以两个 100-token 请求为例(P=96, S0=4, S1=4):

批次 1: req0 = [P, S0]  完整 prefill 100 token,P 的 KV 写入 radix cache
批次 2: req1 = [P, S1] 前缀命中 96 token,只 prefill S1(4 token)

P 的 prefill 只发生一次,重复计算通过跨批次 cache 命中消除。该拆分依赖 cache-aware 策略(lpm 或 dfs-weight),非默认策略,需显式开启。

计算示意图

4.2

 如上图所示,两个请求拆分为两个批次,第一个批次完整 prefill 100 token,P 的 KV 写入 radix cache,第二个批次前缀命中 96 token,只 prefill S1(4 token)。

4.3 方案 3:vLLM 的 Direct Ragged Prefill(同 Batch 前缀去重)

方案 2 将 batch 拆分为两个,会增加第二个批次请求的耗时,是否有更好的思路?我们来看看 vLLM 的方案。vLLM 将同 Batch 内多个请求的 Q 序列进行前缀去重与打平拼接,合并为一次 Direct Ragged Prefill Kernel 调用。利用 FlashInfer 的 Ragged QO 能力与 Paged KV 机制,使共同前缀仅在 Q 中保留一份并执行一次 Prefill 计算,各请求独有后缀通过 Page Table 共享该前缀 KV。

底层原理

以两个 100-token 请求为例,共享 96-token 前缀,后缀各 4-token:

P  = 96-token 共同前缀
S0 = req0 的 4-token 独有后缀
S1 = req1 的 4-token 独有后缀

req0 = [P, S0],长度 100
req1 = [P, S1],长度 100

(1)Q 的处理

利用 FlashInfer BatchPrefillWithPagedKVCacheWrapper 的 ragged QO 能力,将 Q 打平拼接:

Q = [P(96), S0(4), S1(4)]   # 总共 104 个 token
qo_indptr = [0, 100, 104] # 两个 request

request 0: Q = [P, S0], q_len = 100
request 1: Q = [S1], q_len = 4

这样 P 只在 request 0 中计算,request 1 不再重复计算 P。

(2)KV 的处理

request 1 只是不重复计算 P,它的 attention 仍然要看到完整的 [P, S1] KV(100 个 token)。这是通过 Paged KV 的页表(paged_kv_indices)实现的。以 page_size=16 为例,前缀 P 占 6 个完整页(页 0~5),S0、S1 各占 1 页(页 6、页 7),两个请求的 KV 布局如下:

页 0-5:  P 的 96 个 KV token(物理共享,只存一份)
页 6: S0 的 4 个 KV token
页 7: S1 的 4 个 KV token

request 0 的 block table: [0,1,2,3,4,5,6] → 看到 [P, S0],共 100 token
request 1 的 block table: [0,1,2,3,4,5,7] → 看到 [P, S1],共 100 token

paged_kv_indptr = [0, 7, 14]
paged_kv_indices = [0,1,2,3,4,5,6, 0,1,2,3,4,5,7] # 两请求的 block table 拼接
paged_kv_last_page_len = [4, 4]

可以看到 request 1 的 block table 同样包含前缀页 0~5。前缀 KV 页在物理上只存一份,被两个请求的页表共同引用。因此 request 1 的 Q 只有 4 个 token,但通过页表映射,其 attention 覆盖完整的 100 个 KV token,输出与"完整 prefill [P, S1]"完全一致,且 P 的 KV 无需重复计算。

(3)vLLM 中的实现细节

vLLM 里是怎么知道同一个 batch 里有前缀共享的呢?前缀共享的判定依赖 rolling block hash 及其写入时机。vLLM 用每个 block 的内容 hash 标识该 block(rolling block hash:相邻 block 的 hash 存在依赖,前面的 block 一旦变化,后续 hash 全部失效,因此能唯一锁定"以某内容为前缀"的整条序列)。在申请到 block 时就将对应的 hash 值登记进全局哈希表,此时该 block 的 KV 尚未写入,相当于"提早建索引"。于是同一个 batch 内第二个请求调度时,按 hash 命中这些 block,就可以判断命中 cache。由于同 batch 的 prefill forward 分为两步:先计算 QKV 矩阵并将 K/V 写入 KV cache,再执行注意力计算与 softmax,因此同一个批次内执行注意力时,被复用的 KV cache 已经有值,复用是安全的。这是 vLLM 能在同一 batch 内做前缀去重(而非仅跨批次)的根本原因。

对比之下,SGLang 的 radix cache 要等 prefill 完成、KV 写入后才登记进 radix tree,同一 batch 内调度时命中不了尚未写入的 KV。SGLang 当前的 in-batch prefix caching(见 4.2 节)通过把共享前缀请求拆成多个批次、先让第一个请求完整 prefill 生成 KV 再让其余请求复用,确实能实现同 batch 前缀去重,但这种方式受限于"KV 先算好才可命中"的约束,需要额外拆批、引入调度延迟,性能较差。从原理上讲,SGLang 也可以像 vLLM 一样直接复用同 batch 内即将写入的 KV,只是需要较大的改动:把 radix tree 的登记时机提前到 KV 写入之前(与 vLLM 的 rolling block hash 同思路)。

计算示意图

4.3

 如上图所示,同 batch 前缀去重能成立的关键时序是:prefill forward 先执行 QKV 投影并写入 KV cache,再执行注意力读取——因此被共享的前缀 KV 在注意力执行时已有值,复用是安全的。在此基础上,FlashInfer 的 BatchPrefillWithPagedKVCacheWrapper 通过 ragged QO(qo_indptr 打平请求)与 paged KV(多请求共享同一 block table)将去重落到 kernel 层:共享前缀的 Q 只计算一次,多请求通过页表共享同一份前缀 KV。

4.4 方案 4:Cascade Attention

上述 vLLM 的方案也存在缺陷:从 BatchPrefillWithPagedKVCacheWrapper 的 plan 参数可以看到,paged_kv_indices 中存在重复的 block index:共享前缀页在多个请求的 block table 中各出现一次。前缀 KV 在 HBM 里只存一份,但内核中每个请求的 CTA 都会独立扫描完整 KV 范围(含共享前缀页),因此共同前缀会被多次从 HBM 重复读取,在 memory-bound 时这是可观的带宽开销。另外,flash attention 的一个 CTA 只能处理单个请求的 (query 区间, KV 区间) 组合,共享前缀的多个请求无法合并到同一个 CTA 中执行,而必须拆成多个 CTA 分别调度。这样会使计算任务更加分散,增加 CTA 的调度与固定执行开销;尤其在各请求的 query 较短时,单个 CTA 的工作量不足,更难充分利用 GPU 的计算资源。

为了解决这个问题,我们引入了 cascade attention。先计算共同前缀的 prefill,然后再使用 cascade attention:消除共同前缀的 HBM 读取,再对后缀执行 prefill,最后 merge。

底层原理

同样以两个 100-token 请求为例(P=96, S0=4, S1=4)。先看单个 transformer 层内四个阶段的划分(真实模型有 L 层,各阶段的执行层次见下文"执行结构");由于 MultiLevelCascadeAttentionWrapper 要求各 level 引用的 KV cache 已存在,该层流程分为四个阶段:

阶段 A:前缀 causal prefill(独立阶段,可先行完成所有层)
Q = P, KV = P, causal = True
→ 计算 P 的注意力输出,同时生成 P 在各层的 KV cache

阶段 B:cascade shared level
Q = [S0, S1](打包为一个 logical request)
KV = P(读取阶段 A 生成的 cache)
causal = False
→ P 严格位于所有后缀 token 之前,每个后缀 query 都应看到完整 P

阶段 C:cascade unique level
request 0: Q = S0, KV = S0, causal = True
request 1: Q = S1, KV = S1, causal = True
→ 各请求只读自己的后缀(其 KV 已在该层 forward 中实时生成),物理隔离

阶段 D:merge
merge_state_in_place(out, lse, out_i, lse_i)
→ 按 query index 将两级结果做 LSE 加权合并

执行结构:阶段 A 独立跨层先行,阶段 B/C/D 逐层循环。 阶段 A(前缀 causal prefill)是独立阶段,可先一次性跑完所有层:对 P 逐层做 causal prefill(每层用上一层输出的 P hidden state 作为输入,写 P 在该层的 KV cache),得到 P 在所有层的 KV,全程不涉及任何后缀 token。之后才进入阶段 B/C/D——它们封装在 MultiLevelCascadeAttentionWrapper 中,模型循环对每一层调用一次 run(q, kv_cache_at_layer)(注:plan() 只做一次,各层复用同一套辅助结构,见 6.1 节):每层先做该层 S0/S1 的 QKV 投影(实时生成并写入该层后缀 KV),再依次执行 shared(阶段 B)→ unique(阶段 C)→ merge(阶段 D),输出经 FFN 后作为下一层输入。跨层来看,第 L 层的 KV 依赖前 L-1 层的输出,天然逐层产生。

两级结构的 indptr 设计。 阶段 B/C/D 均由 MultiLevelCascadeAttentionWrapper 完成:plan() 一次配置两级结构(指前面的 shared level 和 unique level,实际上本接口支持配置多级),两级通过 qo_indptr_arr 分组;run() 内部依次执行各 level 的 attention 并完成阶段 D 的 LSE merge。只有阶段 A 的前缀 prefill 由独立的 BatchPrefillWithPagedKVCacheWrapper 完成。

qo_indptr_arr[0] = [0, 8]        # shared level: 1 组(S0+S1 共 8 个 query)
qo_indptr_arr[1] = [0, 4, 8] # unique level: 2 组(S0、S1 各自独立)

shared level 把 S0 和 S1 打包为同一组、unique level 各自独立——这正是 cascade 相比方案 3 节省前缀读取的关键机制。

前缀 KV 已存在的场景。 上述流程假设前缀 KV 尚未生成,需要阶段 A 先算。如果前缀 KV 已在 cache 中(如历史批次已计算),可跳过阶段 A,直接对后缀做两级 cascade——这就是第 6 章 decode 场景和 extend prefill 场景的标准用法。

计算示意图

4.4

 如上图所示,整体分为两大阶段,先算公共前缀 prefill,再使用 cascade attention 消除各请求后缀 query 对公共前缀 KV 的重复读取。

4.5 方案 5:掩码合并(kMultiItemScoring / custom_mask)

前面的 Direct Ragged Prefill 方案,通过 qo_indptr 将多个请求剔除重复后参差不齐地拼到一起,而 kv_indices 还是保持多个完整的请求。flashinfer 也提供了另一种更彻底的多请求拼装,即 kMultiItemScoring 和 custom_mask。这两种方案都是将多个请求合并为一个逻辑请求,单次 kernel 内通过 mask 控制各 item 的可见区间,避免拆分为多次 kernel。下面详细介绍 mask 的方案。按 mask 的表达方式,分为两个子实现,属于同一条路线的特例与泛化:

  • kMultiItemScoring:结构化 mask(用少量 metadata 在 kernel 内实时构建),高效,但只支持"公共前缀 + 独立 item"这类结构;
  • custom_mask:用显式 Q×K mask 表达任意可见性,通用,但开销更大。

4.5.1 kMultiItemScoring:结构化高效实现

底层原理

同样以两个 100-token 请求为例,这里将 [P, S0, S1] 编码为一个逻辑 request:

Q  = [P(96), S0(4), S1(4)]   # 总共 104 个 token
KV = [P(96), S0(4), S1(4)]
q_indptr = [0, 104] # 一个 request
kv_indptr = [0, 104]
prefix_len = 96

BatchPrefillWithPagedKVCacheWrapper 接口通过三个 metadata 描述分支结构:

  • prefix_len_ptr:标记前 N 个 token 是公共前缀;
  • token_pos_in_items_ptr:每个后缀 token 在其所属 item 内的位置(1-indexed,0 保留给 delimiter)。前缀 token 不进入该数组——前缀的可见性由 prefix_len_ptr + causal 统一表达,只有 item 内部需要逐 token 位置(判断同 item 的 causal、跨 item 的屏蔽);
  • max_item_len_ptr:所有 item 的最大长度(标量),kernel 用它划分 item 边界。

具体例子(P=96, S0=4, S1=4,逻辑序列 [P(96), S0(4), S1(4)]):

metadata说明
prefix_len_ptr [96] 前 96 个 token 是公共前缀
token_pos_in_items_ptr [1, 2, 3, 4, 1, 2, 3, 4] 只覆盖后缀:S0 的 4 个 token 在 item0 内位置为 1~4;S1 的 4 个 token 在 item1 内重新从 1 开始(pos=1 标记新 item 边界)
max_item_len_ptr [4] 所有 item 的最大长度(S0、S1 均为 4)

kernel 内部基于这些 metadata 构建 branch-aware mask:

Query / KeyPS0S1
P query causal masked masked
S0 query 全可见 causal masked
S1 query 全可见 masked causal

mask 在 QK 计算完成后、softmax 之前逐元素应用:遍历每个 Q-K 对,根据 metadata 判断是否有效,无效则填 -inf,再做 softmax + PV matmul。

计算示意图

4.5.1

 如上图所示,多个请求合并为一个逻辑请求,内核通过掩码控制各 token 的可见区间(跨 item 屏蔽 + item 内 causal)。

4.5.2 custom_mask:任意形状 mask 的通用实现

custom_mask 与 kMultiItemScoring 走同一条"合并为一个逻辑请求"的路线,区别在于用显式 Q×K bool mask 表达可见性:每个 request 是一块 q_len[i] × k_len[i] 的二维 mask,接口内部展平为 1D(总长度 = sum(q_len[i] × k_len[i]))并按 request 分段打包成位图,False 表示该注意力元素被 mask 掉;传入 plan(custom_mask=...) 后 kernel 切换 MaskMode::kCustomcausal 参数被忽略(接口细节见 4.5.3 节)。因此它能表达任意形状的注意力模式(如推测解码的树掩码),不限于"公共前缀 + 独立 item"。

代价是 mask 为显式的 (O(QK)) 数组。在共享前缀场景,单逻辑请求下 mask 大小等于 total × total,其中前缀 causal 部分与跨 item 无效部分占绝大多数,且需要存储/读取完整 mask。其功能更灵活,但通常性能劣于 kMultiItemScoring。

4.5.3 接口:custom_mask / packed_custom_mask

plan() 接受两个参数(二选一;同时提供时 packed_custom_mask 优先,custom_mask 被忽略):

wrapper.plan(
...,
custom_mask=None, # torch.Tensor, dtype=bool, 1D 打平
packed_custom_mask=None, # torch.Tensor, dtype=uint8, segment_packbits 预打包
causal=True, # 当 custom_mask 非 None 时忽略
...
)

custom_masktorch.bool):

  • 1D 打平格式,总长度 = sum(q_len[i] × k_len[i]),按 request 索引、行优先拼接;
  • False 表示对应注意力元素被 mask 掉(填入 -inf);
  • 如果只传 custom_mask 而未预打包,内部会调用 segment_packbits 转换为 uint8 打包格式,有额外开销。

packed_custom_masktorch.uint8):

  • 预打包版本,消除运行时的 segment_packbits 开销;
  • CUDA graph 场景下可在 capture 前预先分配 custom_mask_bufmask_indptr_buf,graph 内部免 sync。

custom_maskpacked_custom_mask 提供时,kernel 内部自动切换到 MaskMode::kCustomcausal 参数被忽略。

4.5.4 典型应用:推测解码的树掩码

custom_mask 最典型的应用场景是推测解码(speculative decoding)的验证阶段。以 SGLang 的 EAGLE 实现为例:

在 verify 阶段,draft 模型产生了一棵有 N 个候选 token 的树(如 32 个 token,按 parent-child 关系形成树结构)。验证时:

  • 前缀部分(已存在 KV):每个 draft token 都应完整看到;
  • draft 之间的注意力:受树结构约束——token 只能看到自己的祖先,不能看到其他分支的 token。

这和标准的 causal mask 不同——causal mask 允许看到所有 draft token,而树掩码额外屏蔽了其他分支。custom_mask 就是用来表达这个树结构的。

与 kMultiItemScoring 的关系

 custom_maskkMultiItemScoring
表达能力 任意 Q×K 可见性 仅"公共前缀 + N 个独立 item"
mask 存储 显式 Q×K 打平数组(sum(q_len × k_len)) 三个 metadata tensor(O(1) 相对序列长度)
开销 高(需存储/读取/打包完整 mask) 低(kernel 内按 metadata 实时计算)
典型场景 推测解码树掩码、任意不规则注意力模式 搜索打分、共享前缀 + 短后缀 prefill

kMultiItemScoring 可以视为 custom_mask 的结构化特例——它用一个枚举类的 mask 模式替代了通用的显式 mask 存储,在 "共享前缀 + 独立分支" 这个约束下用更少的 metadata 达到同样的效果。如果注意力模式无法用 kMultiItemScoring 描述(如树结构、任意稀疏模式),则退回到 custom_mask。两个接口互补:一个是通用方案,一个是结构化高效方案。

5. 性能权衡与选型指南(P×S×M 维度)

5.1 深入算子层:CTA 视角下的方案再分析

上述五种方案,哪种方案性能最好?上述的理论分析并不完备,没有考虑到算子层面的 FlashAttention 内核实现。FlashAttention 内核会做分块,天然带来重复读取,但实际影响多大,还取决于 L2 Cache 的命中情况。下面以 CTA 视角更系统地分析各种方案在 HBM 读取量和 CTA 计算个数方面的对比。

FlashAttention 执行模型

Prefill 注意力内核的实现是 FlashAttention:Q 被切分为多个 Q tile,KV 被切分为多个 KV tile;每个 Q tile 由一个 CTA 负责,该 CTA 按顺序遍历其可见的全部 KV tile(忽略 split-KV 时 Q tile 数 = CTA 数)。因此 KV 天然被重复读取:每个 Q tile 都要把完整 KV 扫描一遍。代码实现上是三层循环:请求层、Q tile 层、KV tile 层,不同请求的 CTA 不会重叠。(split-KV 是可选变体:把 KV 维度也 shard 到多个 CTA,换取每个 CTA 更小的 KV 负载,代价是一次额外 merge;本文忽略。)

一个例子:P=96, S0=S1=4, Q tile=8

五个方案在 CTA 数与 KV 扫描量上的对比(方案 2 依赖跨批次 cache、前缀计算发生在历史批次,不在此单批次口径下对比):

方案Attention CTAKV token 扫描量launch构成
方案 1(朴素) 26 1448 ≈1 两请求各自完整 prefill,连前缀 Q 都重复计算
方案 3(direct ragged) 14 824 ≈1 P 12 tile + S0、S1 各 1 tile,各扫一次 P
方案 4(cascade) 15 728 ≥4 前缀 12 tile + 后缀聚合 1 tile + unique 2 tile + merge
方案 5(kMultiItemScoring) 13 728 ≈1 前缀 12 tile + [S0,S1] 聚合 1 tile

上述方案的核心差异在于后缀 Q 能否聚合到同一个 Q tile:

  • 方案 3:S0、S1 分属不同 request_idx,各占一个独立 CTA、各自完整扫描 P(P 因此多扫 2 次);
  • 方案 5:通过 mask 把 S0、S1 打包进同一 Q tile,P 只被额外扫描 1 次;单次 launch、无 merge,代价是跨 item 的 mask(含无效 QK 计算);
  • 方案 4:通过 shared level 把 S0、S1 打包进同一 Q tile,P 只被额外扫描 1 次;无需 mask,但要 ≥4 次 launch 加 merge kernel。

从这个例子看,方案 3 相比理论最优(KV 扫描量最小的方案 4/5)只多 1 次前缀扫描,差距并不大;且多个 CTA 并发加载同一物理 KV 页时,除首次外大概率命中 L2 cache,实际 HBM 流量远小于"加载次数 × 页大小"的字面计算。但这只不过是一个例子的结论——前缀长度、后缀长度、请求数都会改变相对优劣。下面我们沿着这三个维度(P×S×M)做实测,看看真实的性能对比情况。

5.2 实测评测(P×S×M 笛卡尔积)

为验证上述理论分析,对五种方案在 P(前缀长)× S(后缀长)× M(请求数/batch size) 三维笛卡尔积下做了实测。测试使用 FlashInfer v0.6.15.post1;固定配置为 heads=16/8(GQA)、head_dim=128、page_size=16、fp16,warmup=10、iters=100;三个维度各取两级:P∈{96, 8196}、S∈{4, 256}、M∈{2, 32},共 8 个场景。其中 8196 % 16 = 4,用于覆盖非整页长前缀。

kernel 时间均值(ms,数值越小越好)

方案P96 S4 M2P96 S4 M32P96 S256 M2P96 S256 M32P8196 S4 M2P8196 S4 M32P8196 S256 M2P8196 S256 M32
方案1 朴素 0.037 0.094 0.086 0.437 9.883 145.487 10.469 154.564
方案2 批次拆分+Cache 0.049 0.060 0.137 0.404 5.361 5.909 6.039 14.692
方案3 Direct Ragged Prefill 0.036 0.043 0.113 0.346 5.425 7.422 5.953 14.094
方案4 Cascade Attention 0.114 0.115 0.163 0.505 5.452 5.557 6.312 15.245
方案5a kMultiItemScoring 0.038 0.059 0.131 0.551 5.295 5.448 5.906 14.133
方案5b custom_mask 0.048 0.071 0.159 12.508 12.668 13.007 14.085 51.788

相对方案 1 的加速比(越大越好)

方案P96 S4 M2P96 S4 M32P96 S256 M2P96 S256 M32P8196 S4 M2P8196 S4 M32P8196 S256 M2P8196 S256 M32
方案2 批次拆分+Cache 0.76 1.55 0.63 1.08 1.84 24.62 1.73 10.52
方案3 Direct Ragged Prefill 1.02 2.19 0.76 1.26 1.82 19.60 1.76 10.97
方案4 Cascade Attention 0.33 0.82 0.53 0.86 1.81 26.18 1.66 10.14
方案5a kMultiItemScoring 0.96 1.58 0.66 0.79 1.87 26.70 1.77 10.94
方案5b custom_mask 0.77 1.32 0.54 0.03 0.78 11.19 0.74 2.98

结论

  1. M(共享请求数)是决定性维度。 M=2 时,短前缀场景的优化收益有限且多数为负;长前缀因冗余计算量大,结构化方案仍有约 1.65~1.87x。M=32 时长前缀场景的收益显著放大,最高达到 26.70x。
  2. Direct Ragged Prefill 是最稳健的默认方案。 它在 4 个场景中直接最快,其余多数场景也接近最优;主要例外是"长前缀 + 短后缀 + 大 batch"(P8196_S4_M32),此时 7.422 ms 明显慢于 kMultiItemScoring 的 5.448 ms,后者约快 1.36x。Direct Ragged Prefill 的优势是结构简单、launch 少且不需要额外 metadata/merge。
  3. 长前缀 + 大 batch + 短后缀时,kMultiItemScoring / Cascade Attention / 批次拆分领先。 P8196_S4_M32 下三者分别达到 26.70x、26.18x 和 24.62x,均高于 Direct Ragged Prefill 的 19.60x;前缀极长时,后缀 Q 聚合或只计算后缀带来的收益最明显。其中 kMultiItemScoring 的收益上限受后缀长短约束:它要加载并计算合并后逻辑请求范围内的全部 KV(含不属于当前请求的),无用后缀越长,跨请求的无效 QK 计算浪费越大,因此只在后缀很短时划算。
  4. 短前缀 + 短后缀 + 大 batch 时,Direct Ragged Prefill 明显领先。 P96_S4_M32 下达到 2.19x,高于 kMultiItemScoring 的 1.58x、批次拆分的 1.55x 和 custom_mask 的 1.32x;Cascade Attention 受多次 launch 和 merge 固定开销影响,仅为 0.82x。
  5. Cascade Attention 在短前缀场景普遍落后(P96 下为 0.33~0.86x)。多级 attention、LSE 和 merge 的固定开销在短前缀时抵不过共享收益。
  6. custom_mask 不适合作为共享前缀的常规优化。 显式 Q×K mask 及跨 item 的无效计算在长后缀场景代价极高:P96_S256_M32 仅 0.03x,约比朴素方案慢 29 倍;P8196_S256_M32 也只有 2.98x,明显落后于结构化方案的约 10~11x。
  7. 批次拆分 + Cache 在"长前缀 + 大 batch"下竞争力强。 P8196_S4_M32 达到 24.62x,与 Cascade Attention 和 kMultiItemScoring 接近;P8196_S256_M32 也达到 10.52x。注意这是理想化模拟:KV 已预先构造,未计入真实系统中的 KV append、跨批次调度和排队成本。
  8. P96_S256_M2 下所有优化方案都慢于朴素实现(0.53~0.76x)。请求少、前缀短、后缀长时,朴素 batch 已足够高效,额外的 launch、metadata、mask 或 merge 都成为负收益。

评测局限:本次数据使用 FlashInfer v0.6.15.post1,固定 heads=16/8(GQA)、d=128、page=16、fp16,报告 warmup 后 100 次 CUDA Event 计时的均值。输入、KV cache 和 plan 均在计时前构造,结果不包含 QKV 投影、KV append、页分配、mask 构造和调度成本。GQA 比例、Q tile、缓存冷热、硬件架构和 FlashInfer 版本都会改变绝对数值。

5.3 总结与选型

三种优化方案可以总结为三条路线:

  1. 方案 3:重复加载 KV cache 路线——多个请求还是独立的,公共前缀部分的 KV cache 重复加载;
  2. 方案 5:多余计算路线——公共前缀部分唯一,但后缀部分有多余计算,在 softmax 之前需要 mask 掉;
  3. 方案 4:不重复、不多算的精确计算路线——公共前缀部分不会重复加载,后缀部分也精确按需计算,最后做 merge。

到这里可以回到文章开头提出的第 3 个问题,减少了重复计算或 HBM 读取,性能不一定是最好的。

适用场景速查:

场景特征推荐方案理由
一般默认;少量分支、短后缀 Direct Ragged Prefill launch 最少、无 merge/mask 开销;4 个场景直接最快
前缀长(P≳数千)+ 大 batch(M≫1)+ 后缀短(S≤4) kMultiItemScoring(也可 Cascade Attention / 批次拆分) P8196_S4_M32 下三者达到约 25~27x;后缀各自成 CTA 时 P 被重复加载 32 次,聚合后 P 只扫 1 次。prefill 下不推荐 Cascade Attention:需 ≥4 次 launch,其主战场在 decode 阶段(见第 6 章)
短前缀 + 大 batch(P≤96) Direct Ragged Prefill P96_S4_M32 下 2.19x 明显领先;短前缀的重复扫描命中 L2,几乎免费
追加 prefill(前缀 KV 已存在、后缀较长) Direct Ragged Prefill 后缀长时打包无收益,前缀扫描次数与 Direct Ragged Prefill 相同,merge 成纯开销;跨请求 mask 浪费随后缀长度线性增长
同 batch 共享前缀 KV 尚未生成、后缀短 kMultiItemScoring 单次 kernel 完成全部计算;Cascade Attention 必须先单独 prefill 前缀(≥4 次 launch);Direct Ragged Prefill 也可用但每个 suffix CTA 各扫一次 P
请求数很小(M≤4) 朴素 batch 或 Direct Ragged Prefill 短前缀下多数优化为负收益;长前缀时结构化方案仍可能受益
需要任意形状 mask(如推测解码树掩码,kMultiItemScoring 无法表达) custom_mask 表达能力通用,但共享前缀场景下显式 mask 开销很大

总的来说,Direct Ragged Prefill 与 kMultiItemScoring 加起来可以解决 prefill 阶段的所有场景:常规场景用 Direct Ragged Prefill(通常最优),仅当分支非常多、后缀非常短、重复读取前缀 KV 的 HBM 代价巨大时才改用 kMultiItemScoring。方案 4(Cascade Attention)这条精确计算路线在 prefill 阶段反而平庸,其主战场在 decode 阶段。

6. 场景 3:Decode 阶段的突破:Cascade Attention

Decode 阶段的共享前缀场景与 prefill 有本质不同:每个请求每步只产生 1 个 query token,但前缀和历史 KV 都已存在于 cache 中,此阶段的核心浪费是多个请求重复从 HBM 读取相同的前缀 KV。针对这一浪费,FlashInfer 团队于 2024 年 2 月 2 日在技术博客《Cascade Inference: Memory Bandwidth Efficient Shared Prefix Batch Decoding》中首次提出 Cascade Attention——一种专用于 decode 阶段的 HBM 重复读取优化技术,并建立了工程实现。

6.1 底层原理

问题定义

以 7 个请求共享前缀为例:

P  = 共享前缀 KV(已缓存,如 512 pages)
S0~S6 = 各请求独立的历史 KV(已缓存,长度各异)
Q = 每个请求 1 个新 query token

朴素实现中,7 个请求各自作为独立 CTA 工作项,每个都完整扫描 P——P 被从 HBM 重复读取 7 次。Cascade Attention 的核心思路,就是把"7 个 query 各扫一遍 P"拆成"先一起扫一遍 P、再各自扫自己的后缀"两步,从而消除这种重复读取。

两级结构

FlashInfer 通过 MultiLevelCascadeAttentionWrapper 将该问题拆分为两级:

Level 0(shared level):
Q = 7 个 query 打包为一个 logical request
KV = P
causal = False(P 是历史 token,所有 query 都应完整看到)
→ P 只被扫描一次(7 个 query 聚合到同一个 Q tile)

Level 1(unique level):
request i: Q = 1 个 query, KV = S_i
causal = True
→ 各请求只读自己的后缀 KV,物理隔离,零跨请求扫描

Merge:
merge_state_in_place(out, lse, out_i, lse_i)
按 query index 将两级输出合并(LSE 加权的在线 softmax 合并)

对应的 plan() 数据结构:

qo_indptr_arr       = [[0, 7], [0, 1, 2, ..., 7]]
↑ shared: 1 组 ↑ unique: 7 组
paged_kv_indptr_arr = [shared_indptr, unique_indptr]
paged_kv_indices_arr= [shared_indices, unique_indices]

plan() 设置一次后,32 层 transformer 可复用同一套辅助结构,每层只调 run(q, kv_cache_at_layer[i])

上面的两级结构是否始终有效?特别是各请求的后缀是长期积累的长历史 KV 时,这种拆分会不会失效?这正是下一节要回答的问题。

6.2 关键疑问:后缀很长时 cascade 还有收益吗?

一个自然的疑问是:既然 prefill 阶段"后缀长则 Q 打包无收益",那么 decode 时各请求的后缀是长期积累的长历史 KV,cascade 是不是也失效了?普通 batch decode 就够了?

答案是否定的——decode 下 cascade 的收益与后缀长短无关。关键在于收益来源不同:

  • prefill 中方案 4/5 的收益来自"把多个请求的后缀 query 合并到同一 Q tile",后缀一长,每个请求的后缀自己就占满多个 tile,合并不再减少 tile 数;
  • decode 中每个请求永远只有 1 个 query token,cascade 的收益来自 Level 0 把 N 个 query 打包进同一个 Q tile 去扫 P,而不是合并后缀。

对比两者的执行过程:

普通 batch decode:
request 0: 1 query × [P, S0] → 扫一遍 P
request 1: 1 query × [P, S1] → 又扫一遍 P
...
request 6: 1 query × [P, S6] → 再扫一遍 P
P 共被扫 7 次

Cascade:
Level 0: [Q0..Q6] 7 个 query 打包为 1 个 logical request
→ 1 个 Q tile 扫 P,P 只扫 1 次
Level 1: 各 query 扫各自的 S_i(与普通 batch 完全相同)

两点关键认识:

  1. Level 0 的打包与后缀长短无关。 query 数量永远是 N 个(每请求 1 个),无论 S_i 是 10 还是 10000,N 个 query 都能被紧密打包进 Q tile——在 N 不超过 Q tile packed-row 容量时落进同一个 tile,N 更大时按 ceil(N·g/T_Q) 切分为少数几个 tile,收益随 tile 数线性递减但方向不变。后缀长度本身完全不影响 shared level 的聚合效果。
  2. Level 1 的后缀扫描量与普通 batch 完全一致。 每个 query 必须读自己的历史 KV,这部分不可避免,也不因 cascade 而改变。cascade 没有引入任何额外扫描。

因此 decode 下 cascade 节省的 "(N-1) 次 P 扫描"是与后缀长度解耦的。真正决定收益大小的是 P 的长度和请求数 N:P 越长、N 越大,节省越多。后缀极长时,P 的节省在总工作量中的占比会下降,但由于 decode 的 merge 只涉及 1 token、开销几乎为零,这笔节省始终是零成本的纯收益——不存在"后缀太长所以不值得用 cascade"的情形。

此外,cascade 对 decode 还有一层经常被忽略的收益:普通 batch decode 中每个请求的 1 个 query 各自为一个 CTA,每个 CTA 的工作量极小(单 token × 长 KV),GPU 计算利用率很低。Cascade 的 shared level 将 N 个 decode query 合并到一个 Q tile/CTA,等效于将原本分散的 N 个"单 query × 长 KV"矩阵乘法合并为一个"多 query × 同一 KV"的大 batch 矩阵乘法。这不仅减少了 KV 的重复读取,还让每个 CTA 内的计算密度更大,更好地利用 GPU 计算单元(如 Tensor Core),在 decode 这一 memory-bound 场景下可能进一步缓解算力闲置。

这也正体现了 Level 1 "分层隔离"的意义:各请求的后缀 KV 分为各自独立的 logical request,各 query 只读自己的后缀,物理上杜绝了跨请求扫描(这正是 kMultiItemScoring 无法做到的,见 6.3)。

6.3 为什么 decode 不用 kMultiItemScoring

理论上 kMultiItemScoring 也能处理该场景,但实践中有两个障碍:

(1)跨请求 KV 扫描浪费大。 decode 时各请求的后缀是长期积累的历史 KV,而 kMultiItemScoring 要求每个 query 遍历全部 KV([P, S0, S1, ..., S6]),跨请求的 KV 在 QK 计算后被 mask 掉。后缀越长,浪费越严重:

kMultiItemScoring 每个 query 的 KV 扫描: P + S0 + S1 + ... + S6
其中有效的只有: P + S_i
浪费比例: 6/7 的后缀扫描全部浪费

Cascade 的 unique level 则让每个请求只读自己的后缀,物理上避免了这种浪费。

(2)flashinfer API 未暴露。 kMultiItemScoring 通过 prefix_len_ptr 启用,而该参数只在 BatchPrefillWithPagedKVCacheWrapper.plan() 中暴露;BatchDecodeWithPagedKVCacheWrapper 的公开接口没有这些参数。

7. 行业实践与业务落地

7.1 搜索打分场景

搜索 query 召回多个文档,在展示给用户之前需要结合用户信息、历史 query 等当前信息做打分排序。使用生成式模型做打分排序时,经常是 prompt = common prompt + doc 算一个 doc 的打分,一共有几十个打分对应到几十个文档。这些请求通常组成一个大 batch,符合第 2 章提到的第二类场景:在一个 batch 里有公共前缀。前文的分析表明,prefill 阶段由 Direct Ragged Prefill 与 kMultiItemScoring 两个方案即可覆盖全部场景

  • Direct Ragged Prefill:让重复前缀只计算一次 prefill,但重复部分的 prefill 的 KV cache 会多次读取;
  • kMultiItemScoring:让重复前缀只计算一次 prefill,单次 kernel 完成;doc 较短时收益最大,doc 变长后跨 item 的无效计算和 mask 浪费线性增长。

值得一提的是,kMultiItemScoring 的名字正来源于此场景——"多 item 打分"(一个公共 prompt + N 个 doc 分别打分)就是这个 mask 模式的设计初衷。考虑到搜索场景下 doc 不可能很短,Direct Ragged Prefill 通常是最佳方案——这是常用来解决同一个 batch 里有公共前缀的最朴素最直接的优化方案,kMultiItemScoring 是替代方案。

7.2 Beam Search

Beam search 是 decode 阶段共享前缀的典型场景:同一 prompt 的 N 个 beam 共享全部历史 KV,每个 beam 每步只新增 1 个 token。N 个 beam 作为同批次请求,前缀部分通过 cascade shared level 一次扫描,beam 各自的 KV 通过 unique level 隔离。

decode 阶段使用 cascade attention 减少 HBM 重复读取,但考虑到多次 kernel launch 和 merge 计算的代价,通常在公共前缀比较长(譬如 1024)时才有收益。

当前我在 SGLang 上实现的 beam search(https://github.com/cswuyg/sglang/tree/feature/beam\_search\_update\_0801)没有直接使用 MultiLevelCascadeAttentionWrapper,而是手动组织多阶段调用——原因是其后缀阶段希望使用 BatchDecodeWithPagedKVCacheWrapper(q_len=1 的 decode 优化 kernel),而 MultiLevelCascadeAttentionWrapper 所有 level 统一使用 BatchPrefillWithPagedKVCacheWrapper

7.3 学术前沿:与 Cascade 思路一致的行业工作

以下工作围绕"共享前缀分解为 shared + unique 两级"的思路展开,该思路在 prefill(第 4 章方案 4:Cascade Attention)与 decode(第 6 章)阶段均有对应实现。主要信息来自:https://www.zhihu.com/question/385229505/answer/3602332966

  1. 《Cascade Inference: Memory Bandwidth Efficient Shared Prefix Batch Decoding》(FlashInfer blog)Cascade Attention 的原始出处。核心思想是将多个共享前缀请求的 attention 拆分为两部分:先对所有请求的 query 计算共享前缀的 attention(shared level),再对各请求单独计算独有后缀的 attention(unique level),最后通过 log-sum-exp 合并两级 partial attention state。FlashInfer 的 MultiLevelCascadeAttentionWrapper 是该思想的工程实现。
  2. 《RelayAttention for Efficient Large Language Model Serving with Long System Prompts》(arXiv:2402.14808)与 Cascade Attention 几乎相同的思路:将长 system prompt(共享前缀)与各请求的对话历史分离计算,通过两次 attention pass("relay" shared prefix attention + per-request unique attention)减少重复的 prefix KV 读取。同样适用于请求共享长 system prompt 的 decode 场景。
  3. 《Hydragen: High-Throughput LLM Inference with Shared Prefixes》(arXiv:2402.11599)同样采用将共享前缀注意力分解为 "prefix attention" 和 "unique suffix attention" 的方案。相比 Cascade Attention 和 RelayAttention,Hydragen 观察到一个额外好处:合并后的 prefix attention 实质是"多 query × 同一 prefix KV"的大 batch 矩阵乘法,相比分散的小 batch 有更高的 GPU 计算利用率,在 decode 场景下可能从 memory-bound 转向 compute-bound。
  4. 《ChunkAttention: Efficient Self-Attention with Prefix-Aware KV Cache and Two-Phase Partition》(arXiv:2402.15220)提出 ChunkAttention,包含两个核心设计:(1)基于前缀树的 PAKV(Prefix-Aware KV)存储,运行时自动复用多个请求共享前缀的 KV cache,降低 KV 显存占用;(2)面向 decode 的两阶段分区注意力内核——先合并同批次共享前缀的 query 统一计算共享前缀 attention(减少 HBM 访存、利用大 batch 发挥 Tensor Core 算力),再逐请求计算独有后缀 attention 并融合分片结果。核心思路与 Cascade Attention 一致。

注:本文也发表于知乎和公众号

← 返回列表