KV Cache 与 PagedAttention:显存复用与长上下文推理

KV Cache 是自回归推理最大的显存开销来源,PagedAttention 用操作系统页式管理思想把显存复用做到极致,配合 KV 量化、剪枝与前缀共享,共同支撑长上下文推理。

自回归推理每生成一个 token,都要把前文每个位置的 Key/Value 缓存下来,这就是 KV Cache。随着上下文变长,KV Cache 的显存占用线性增长,很快成为比模型权重更大的显存黑洞——这也是「长上下文」在推理侧远比训练侧难的根本原因。本文讲透 KV Cache 的显存公式、PagedAttention 的页式复用、KV 量化/剪枝/前缀共享三大降本手段,以及长上下文推理的工程策略与调优参数。

前置:/ai-vllm-system/(连续批处理与内存高效推理)、/ai-attention-optimization/(注意力计算与 KV Cache 优化)、/ai-cuda-memory-optimization/(GPU 显存管理与复用)、/ai-llm-quantization/(量化压缩)。

目录

1. KV Cache:自回归推理的显存黑洞

KV Cache 是什么:
□ 解码时每个 token 的注意力要复用历史 Key/Value
□ 缓存下来 → 避免每步重复计算前文注意力
□ 每层、每个注意力头各存一份

显存公式:
KV = 2 × layers × kv_heads × head_dim × seq_len × bytes
□ 2:Key 与 Value 各一份
□ kv_heads:GQA 后的 KV 头数
□ seq_len:上下文长度;bytes:精度(FP16=2,INT8=1)
def kv_cache_size(layers, kv_heads, head_dim, seq_len, bytes_per=2):
    return 2 * layers * kv_heads * head_dim * seq_len * bytes_per

# Llama-2-70B:80 层,GQA 8 个 KV 头,head_dim=128
for seq in (2048, 8192, 32768):
    gb = kv_cache_size(80, 8, 128, seq) / 1024**3
    print(f"seq_len={seq:>6} → {gb:.2f} GiB")
# 2048 → 0.63 GiB;8192 → 2.50 GiB;32768 → 10.00 GiB
# batch=8、seq=32K → 80 GiB,直接吃掉整张 A100

工程要点:KV Cache 的显存随「上下文长度 × 层数 × 头数 × 精度」线性增长,长上下文场景里它很快超过模型权重本身——这是所有 KV 优化技术(分页、量化、剪枝、前缀共享)共同的出发点。

2. PagedAttention 的核心思想:页式管理

传统框架按「最大序列长度」为每条请求预分配连续张量,碎片浪费惊人。PagedAttention 直接借用了操作系统的虚拟内存分页思想。

传统方案的问题:
□ 预分配 (max_batch, max_seq_len) 的连续 KV Cache
□ 100 token 与 4000 token 的请求占一样大的张量
  → 内部碎片吃一半以上显存
□ batch 容量由「显存」决定,而不是由「算力」决定

PagedAttention 的设计:
□ KV Cache 切成固定块(默认 16 个 token)
□ 逻辑连续的序列 → 物理上存放在不连续的块里
□ 每请求一张块表(block table):逻辑位置 → 物理块
□ 空闲块进全局池,任何请求按需分配
□ 内部碎片只剩「最后一个块」,块池共享消除外部碎片
物理布局示意:
请求 A(35 token)→ 块 7、块 12、块 3(末块只用 3/16)
请求 B(10 token) → 块 5
请求 C(70 token) → 块 8、块 1、块 2、块 0、块 9
块表(A):[逻辑 0-15]→块 7、[16-31]→块 12、[32-34]→块 3

工程要点:PagedAttention 的本质是「把连续的 KV Cache 变成不连续的页式分配」——按需分配块、尾部块才有碎片、块池全局复用。它把显存利用率从「按最大值预分配」提升到「按实际长度分配」,是 vLLM 吞吐领先的基石。

3. vLLM 的实现:块表、共享与写时复制

vLLM 把页式思想落成了一套完整运行时:块分配器、引用计数和写时复制。

三大运行时组件:
□ Block Allocator:显存池按块管理,支持换出到 CPU RAM
□ Block Table:每请求一张,decode 时按块加载 KV
□ Reference Count + 写时复制(COW):
  - 物理块可被多请求共享 → 计数递增
  - 写共享块时先复制一份私有块

共享场景的价值:
□ 同一 prompt 并行采样 4 条答案 → 前缀 KV 完全共享
□ beam search:分叉前共享,分叉时 COW
□ RAG/多轮:system prompt 被大量请求复用
class Block:
    def __init__(self, bid):
        self.bid, self.refcount = bid, 1

def copy_on_write(seq, logical_idx):
    old = seq.block_table[logical_idx]
    if old.refcount > 1:        # 有人共享 → 复制一份
        new = alloc_block()
        new.data.copy_(old.data)
        old.refcount -= 1
        seq.block_table[logical_idx] = new.bid

工程要点:vLLM 的关键是「块 + 引用计数 + COW」三位一体——块表让非连续存储可寻址,引用计数让共享显存可复用,写时复制让共享块安全被写。实测 beam_width=4、seq=2048 时共享可省 60%-70% 的重复 KV Cache。

4. KV 量化:INT8/FP8 压缩显存

量化把 KV Cache 从 FP16 压到更低精度,直接按比例省显存,省下的带宽换更大 batch 或更长上下文。

量化原理:
□ FP16 → INT8:显存减半(2 字节 → 1 字节),主流
□ FP16 → FP8(E4M3):减半,精度略好于 INT8
□ 常见方案:per-channel / per-token 缩放
□ 长上下文下累积误差放大 → 敏感层保留 FP16(混合精度)

收益评估:
□ 单请求 seq=32K、80 层:FP16≈10 GiB → INT8≈5 GiB
□ 省下的显存直接换成更大 batch 或更长 max_model_len
□ 代价:1%-2% 困惑度上升(长上下文更明显)
# vLLM 开启 KV Cache FP8 量化,配合长上下文
python -m vllm.entrypoints.openai.api_server \
    --model meta-llama/Meta-Llama-3-8B-Instruct \
    --kv-cache-dtype fp8 \
    --max-model-len 32768

工程要点:KV 量化是「最简单粗暴的显存减半」——INT8/FP8 直接压掉一半,省下的显存换成 batch 或上下文长度。但长上下文下累积误差会放大,敏感层应保留 FP16,上线前务必做端到端质量回归。

5. KV 剪枝与淘汰:不重要的 token 不缓存

不是所有历史 token 都值得缓存。KV 剪枝的思路是:把「不再重要的位置」从缓存里淘汰掉。

为什么可以剪:
□ 注意力天然稀疏:多数 token 只 attend 到少数关键位置
□ 长上下文中大量是噪音(重复文本、无关段落)

两类经典方法:
□ H2O(Heavy Hitter Oracle):
  - 保留累计注意力分数最高的 token
□ StreamingLLM:
  - 保留初始 token(attention sink)保证数值稳定
  - 滑动窗口保留最近 token,中间直接丢弃
  - 显存占用与上下文长度「脱钩」,O(seq) → O(budget)

组合:初始 token + 最近窗口 + 高注意力 token
sink_size, window_size = 4, 256
h2o_size = 512 - sink_size - window_size

def should_cache(pos, scores, is_recent):
    if pos < sink_size or is_recent:
        return True                  # 初始/最近:必缓存
    return pos in scores.argsort()[-h2o_size:]  # 高注意力

工程要点:KV 剪枝的核心是「把 O(seq_len) 的缓存压到 O(budget)」——StreamingLLM 用「初始 token + 滑动窗口」保证数值稳定,H2O 用「高注意力 token」保证质量。显存从随上下文线性增长变成「封顶」,是 100K+ 超长上下文的关键手段,但要接受偶发质量回退。

6. 前缀共享与 Prefix Caching:复用历史计算

RAG、Agent、多轮对话里大量请求共享同一段 prompt 前缀。把这段前缀的 KV 算一次、复用多次,性价比最高。

共享场景:
□ System prompt:几乎所有请求共享同一段指令
□ RAG 上下文:同一文档被多次检索提问
□ 多轮对话 / 并行采样:历史 KV 天然可复用

两种形态:
□ vLLM Prefix Caching(COW 物理共享):
  - 请求到达时按 token 前缀匹配缓存块,命中则跳过 prefill
□ SGLang RadixAttention(结构共享):
  - 前缀组织成前缀树,任意前缀细粒度共享

收益:
□ 共享前缀场景 TTFT 降 50%-90%
□ 省掉的是 prefill 计算 → 相当于白赚一批算力
# vLLM 开启前缀缓存(RAG/Agent 场景强烈建议)
python -m vllm.entrypoints.openai.api_server \
    --model ... \
    --enable-prefix-caching \
    --max-model-len 32768
命中率优化:
□ 把高复用前缀放在 prompt 开头(缓存按前缀匹配)
□ 保持 system prompt 完全一致,避免拼接顺序抖动

工程要点:前缀共享是「把重复 prefill 直接干掉」——vLLM 的 COW 物理共享简单有效,SGLang 的 Radix Tree 共享更细。RAG/Agent 场景强烈建议开启 prefix caching,实测共享前缀请求的 TTFT 可降 50%-90%。

7. 长上下文推理策略:显存、精度与算法

长上下文(32K-1M tokens)推理是把前面所有手段组合起来的关键战场。

显存侧组合拳:
□ KV 量化:INT8/FP8 减半
□ 剪枝/淘汰:StreamingLLM 封顶缓存
□ 前缀共享:跨请求省显存
□ 分级存储:热 KV 在 GPU,冷 KV 在 CPU/NVMe(offload)

精度侧:混合精度(敏感层 FP16)、RoPE 大角度数值稳定性

算法侧:
□ 稀疏/近似注意力:局部窗口 + 全局 token
□ 上下文压缩:先摘要、再推理(长文档分段问答)

工程侧:
□ max_model_len 与 batch 容量互为代价
□ 长请求与短请求混合调度(chunked prefill 交错)
容量规划示例(单卡 A100 80GB,FP8 KV):
权重(8B FP16)≈16 GiB + 激活/工作区 ≈10 GiB
剩余 ≈54 GiB 全给 KV → seq=32K 单请求 KV≈0.5 GiB
→ 理论并发 ~100 条,实际受调度/共享/碎片影响需压测

工程要点:长上下文不是单一技术能解决的,而是「量化 × 剪枝 × 共享 × 分层存储」的组合拳,外加稀疏注意力和摘要压缩。显存规划「先留足权重与激活,剩余全给 KV」,再用压测校准真实并发容量。

8. 调优参数与实践:从单卡到集群

把 KV 相关调优参数落到真实引擎上,给出可直接上手的建议。

vLLM 核心参数:
□ --gpu-memory-utilization:0.85-0.95
  - 太高:kernel 工作区不足,长序列 OOM
  - 太低:KV 池小,并发受限
□ --max-model-len:单请求上下文上限,与并发互为代价
□ --max-num-seqs:最大并发请求数(防 OOM 保险丝)
□ --enable-prefix-caching:共享前缀场景必开
□ --kv-cache-dtype:fp8/int8,长上下文推荐
□ --enable-chunked-prefill:长 prompt 与短请求混合时开

监控:/metrics 暴露 free/used blocks、前缀命中率、抢占次数
# 生产级启动:长上下文 + 前缀缓存 + FP8 KV
python -m vllm.entrypoints.openai.api_server \
    --model meta-llama/Meta-Llama-3-8B-Instruct \
    --tensor-parallel-size 4 \
    --gpu-memory-utilization 0.92 \
    --max-model-len 65536 \
    --kv-cache-dtype fp8 \
    --enable-prefix-caching \
    --enable-chunked-prefill \
    --max-num-seqs 256

工程要点:KV 调优是「一个显存预算的分配游戏」——权重固定、激活留足、其余全进 KV 池;参数互相牵制,必须以「业务上下文长度分布」为输入压测校准,再定配置。集群侧 TP/PP 让 KV 池随卡数线性增长,DP 靠 session 亲和保住前缀命中。

9. 局限性与踩坑

KV 优化技术各有边界,以下是高频踩坑清单。

局限性:
□ 页式管理:极短序列(seq<64)时块表/分配开销可能抵消收益
□ KV 量化:长上下文 + 低精度 → 质量明显回退的案例常见
□ 剪枝:高信息密度任务(代码、数学)剪枝损失大
□ 前缀共享:请求前缀不一致(用户 prompt 在前)时命中率极低

高频踩坑:
□ gpu-memory-utilization 设 0.98 → 长序列请求偶发 OOM
□ 量化后没做端到端质量回归 → 生成悄悄变差
□ 开启 prefix caching 但拼接顺序抖动 → 命中率为零
□ 长上下文 + batch=1 单条大请求占满显存 → 无并发能力

选型建议:
□ 在线服务:PagedAttention + FP8 KV + prefix caching
□ 超长上下文:加剪枝/StreamingLLM 或摘要压缩
□ 极短交互:评估页式开销,短序列直接预分配

工程要点:KV 优化的边界在于「场景」——页式管理不适合极短序列,量化要质量回归兜底,前缀共享要 prompt 布局配合。上线前用「真实长度分布 + 混合调度」压测,紧盯 KV 池水位与命中率,避免优化变负优化。

10. 速查表与一句话记忆

问题一句话答案
KV Cache 是什么解码时缓存的历史 Key/Value,避免重复计算前文注意力
显存怎么算2 × 层数 × KV 头 × head_dim × 序列长度 × 字节数
为什么是显存黑洞随上下文线性增长,长上下文下超过权重
PagedAttention 解决什么按需分块分配,消除预分配的内部碎片
怎么做到显存共享引用计数 + 写时复制(COW)
KV 量化降多少INT8/FP8 直接减半
长上下文怎么封顶剪枝/淘汰(StreamingLLM、H2O)
共享前缀省什么重复 prefill 计算,TTFT 降 50%-90%
三大降本手段量化、剪枝、前缀共享
最该看的监控KV 池 free blocks、抢占次数、前缀命中率

一句话记忆:KV Cache = 随上下文线性增长的显存黑洞(公式 2×层×头×维×长×精度);PagedAttention = 页式块表 + 引用计数 + COW 复用;三大降本 = 量化减半(INT8/FP8)、剪枝封顶(StreamingLLM/H2O)、前缀共享白赚(TTFT 降 90%)——长上下文推理是「显存预算分配 + 组合拳」的游戏。

延伸阅读

  • /ai-vllm-system/ — vLLM 的连续批处理与 PagedAttention 实现
  • /ai-attention-optimization/ — 注意力计算与 KV Cache 优化
  • /ai-cuda-memory-optimization/ — GPU 显存管理与碎片优化
  • /ai-llm-quantization/ — 权重与 KV 量化压缩
  • /ai-kernel-fusion-optimization/ — 推理内核与算子融合
  • 高性能计算专题 — GPU 内核与显存优化

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. AI 推理安全:内容安全、提示注入与模型防护
  2. LLM 服务可观测性:吞吐、时延、token 与成本监控
  3. FlashAttention 与高效注意力内核:IO 感知与分块计算