Tensor NotesNotes on machine learning
Serving

KV Cache 的显存账本:PagedAttention 与连续批处理

KV Cache 的显存账本:PagedAttention 与连续批处理题图

推理服务能扛多少并发,几乎完全由 KV Cache 决定。很多团队在调度器上花大力气,却没先把这笔显存账算清楚。这篇先算账,再谈优化。

一、把账算出来

每生成一个 token,每一层都要缓存它的 K 和 V。单个 token 在整个模型上占用的显存是:

bytes/token = 2 × nlayers × nkv_heads × dhead × precision

系数 2 是 K 和 V 各一份。代入一个 70B 规模的典型配置(80 层,8 个 KV 头,头维度 128,FP16):

2 × 80 × 8 × 128 × 2 bytes = 327,680 bytes ≈ 320 KB / token

一条 4096 token 的对话就是 1.28 GB。80 GB 的卡装完 140 GB 的模型权重(假设已经切到 2 张卡,每卡 70 GB)之后,剩给 KV Cache 的可能只有几 GB——也就是三四条并发。这就是为什么长上下文服务这么贵。

显存预算分解与分页块分配示意
图 1 · 显存构成。碎片那一块在朴素实现里可以占到 30% 以上

二、分组查询注意力:最便宜的一刀

上面公式里的 nkv_heads 是可以动的。多查询注意力(MQA)把所有查询头共享一组 KV,分组查询注意力(GQA)是折中:每 g 个查询头共享一组。

方案KV 头数(Q 头 = 64)每 token 显存质量影响
MHA642560 KB基准
GQA-88320 KB−0.3% 左右
MQA140 KB−1.5%,长文本上更明显

GQA-8 是目前的事实标准:显存降到 1/8,质量损失在噪声范围内。这个决策必须在训练时做,事后无法转换(有从 MHA 蒸馏到 GQA 的方法,但需要额外训练)。如果你正在规划一个要长期服务的模型,这是最重要的架构决策之一。

三、碎片:朴素实现浪费掉的那部分

传统做法是为每个请求预分配一段连续显存,大小按 max_seq_len 算。问题显而易见:

  • 内部碎片:请求实际只用了 300 token,却按 4096 预留。浪费 92%。
  • 外部碎片:请求结束释放出零散空间,新请求要连续大块,塞不进去。
  • 无法共享:同一个系统提示被 100 个请求复用,却存了 100 份。

实测下来,朴素分配器的有效利用率通常只有 25%–40%。也就是说大部分显存是被浪费掉的,而不是被用掉的。

四、PagedAttention:借操作系统的思路

核心想法直接来自虚拟内存:把 KV Cache 切成固定大小的块(典型 16 个 token 一块),维护一张块表把逻辑位置映射到物理块。序列在逻辑上连续,在物理上可以散落各处。

# 概念示意
block_table[seq_id] = [7, 3, 19, 42]   # 逻辑块 → 物理块号

def locate(seq_id, pos, block_size=16):
    blk = block_table[seq_id][pos // block_size]
    return blk * block_size + pos % block_size

带来三个直接收益:

  1. 内部碎片降到不超过一个块。最坏情况浪费 15 个 token 的空间,相对 4096 的预留可以忽略。
  2. 外部碎片消失。所有块等大,任何空闲块都能用。
  3. 写时复制共享。相同前缀的请求指向同一批物理块,只在分叉时才复制。系统提示、few-shot 示例、并行采样的多个候选,全都能受益。

前缀共享的收益在实际业务里往往最大。我们的场景里系统提示约 1200 token,并发 64 时,共享前后 KV Cache 占用从 24 GB 降到 0.4 GB。

块大小的权衡块越小碎片越少,但块表越大、内核里的间接寻址开销越高。16 是常见默认值;如果你的请求普遍很短(比如平均 100 token 的分类任务),8 更划算;长文档场景可以调到 32。

五、连续批处理

静态批处理要求整批一起开始、一起结束。批里有一条请求要生成 2000 token,其余生成 50 token 就完事了的请求也得陪着占坑。GPU 利用率被最长的那条拖死。

连续批处理(也叫迭代级调度)把调度粒度从"一批"降到"一步":每生成一个 token 就重新评估一次,完成的请求立刻释放,等待队列里的新请求立刻补位。

在我们的线上负载(请求长度方差很大)上,切换到连续批处理后吞吐提升了 3.4 倍,P50 延迟基本不变,P99 延迟下降 40%。这是投入产出比最高的一项优化,如果还在用静态批处理,先做这个。

六、抢占策略

显存不够时必须踢掉某些请求。两种回收方式:

  • 换出(swap):把 KV 块搬到主机内存,恢复时搬回来。PCIe 4.0 x16 大约 25 GB/s,搬 1 GB 需要 40 ms。
  • 重算(recompute):直接丢弃,恢复时重跑一次 prefill。prefill 是计算密集的,但高度并行,实际常常比换出更快。

经验规则:序列较短(< 1000 token)时重算更优;序列很长时换出更优,因为 prefill 的代价随长度平方增长。可以按长度设一个阈值动态选择。

抢占谁也需要策略。纯 FIFO 会让长请求饿死;纯 LIFO 会让先来的用户体验极差。我们用的是「优先抢占剩余生成量估计最大的请求」,配合一个逐渐增长的优先级防止饿死。

七、量化 KV Cache

最后一招是把 KV 本身量化到 INT8 甚至 FP8。显存直接减半,质量损失在多数任务上小于 0.5 个点。几个注意点:

  • K 比 V 更敏感,可以考虑 K 用 FP8、V 用 INT8 的混合方案。
  • 按 head 分组算 scale,不要 per-tensor。
  • 注意力汇聚(sink)token 建议保持原精度,它们的数值幅度和其余 token 差异很大,一起量化会把 scale 拉坏。

把这一节和第二节的 GQA 叠加,相对朴素 MHA + FP16 的方案,每 token 显存可以降到 1/16。并发上限跟着涨一个数量级,这才是长上下文服务能做到可负担的原因。

下一篇:Agent 的工具调用:状态、重试与失败边界…

继续阅读

Related