news 2026/9/27 5:45:17

KV Cache 原理与工程实践:大模型推理加速与显存优化

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
KV Cache 原理与工程实践:大模型推理加速与显存优化

1. 从一次推理延迟排查说起:KV Cache 到底是什么

前阵子帮朋友看一个文本生成的推理服务,现象很典型:单条请求响应挺快,一旦并发上来,延迟直接飙到没法用,GPU 显存还涨得厉害。把 profiling 打开一看,绝大部分时间耗在注意力计算上,而且每生成一个 token,前面所有 token 的键和值都被重新算了一遍。问题就出在这里——KV Cache(Key-Value Cache)没做好,或者说根本没意识到它在自回归生成里的分量。

先把结论摆在前面:KV Cache 是大模型自回归推理阶段的一种缓存机制。它把每一层注意力里已经算过的 Key 和 Value 张量存下来,生成新 token 时直接复用,不再重复计算历史部分。就这么一个动作,能把生成第 n 个 token 的注意力计算量从 O(n²) 降到 O(n),实际推理速度提升往往是几倍到十几倍。代价是显存占用随序列长度线性增长,这也是为什么长上下文场景下显存总是紧张。

这篇内容适合三类人看:一是刚接触 Transformer 推理、搞不清 prefill 和 decode 区别的入门者;二是正在做推理服务、被延迟和显存两头夹击的工程同学;三是想搞明白“为什么大模型吐字速度会越来越慢”的普通使用者。我会从注意力机制的基本计算讲起,把 KV Cache 为什么存在、怎么算、显存怎么估、坑在哪里,一层层拆开。涉及到的参数计算我会给出具体数字,代码部分用 PyTorch 风格示意,方便你直接对照自己的实现。

需要提前说明的是,下面关于工程实现的部分,比如分页管理、量化缓存这些,是基于当前主流推理框架的常见做法做的合理补充,不是某一家的私有方案,你可以按自己用的框架去对应。

2. 为什么需要 KV Cache:自回归生成的重复计算问题

2.1 自回归生成的基本流程

要理解 KV Cache,得先接受一个前提:现在主流的大语言模型都是自回归(autoregressive)生成的。也就是说,模型一次只吐一个 token,然后把新吐出来的 token 接到输入后面,再预测下一个,如此循环。你看到的一段几百字的回答,背后是几百次前向传播。

每一次前向传播,输入都是“原始 prompt + 已经生成的所有 token”。假设 prompt 有 100 个 token,已经生成了 50 个,那第 51 次前向传播的输入长度就是 150。注意,这 150 个 token 里,前 149 个在上一轮已经算过了,只有最后一个是新的。

问题来了:如果每一轮都把 150 个 token 完整地过一遍注意力,那前 149 个 token 的 Key 和 Value 就被反复计算了 51 次。这就是纯粹的浪费。序列越长,浪费越夸张。

2.2 注意力机制里到底算了什么

我们回顾一下缩放点积注意力。给定输入 X,通过三个线性投影得到 Query、Key、Value:

Q = X @ W_q K = X @ W_k V = X @ W_v

然后注意力输出是:

Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V

关键在于,对于第 t 个位置的 token,它要算注意力,需要的是所有位置(1 到 t)的 K 和 V,以及自己这个位置的 Q。它不需要历史位置的 Q,因为 Q 只用来和 K 做匹配,每个位置各算各的。

所以,历史 token 的 Q 算完就可以扔了,但历史 token 的 K 和 V 后面每一步都要用。这就是 KV Cache 名字的由来——只缓存 K 和 V。

2.3 不做缓存时的计算量

假设序列长度为 n,隐藏维度为 d,层数为 L。单层注意力里,Q、K、V 的投影各是 O(n·d²),注意力矩阵 Q@K^T 是 O(n²·d)。生成第 n 个 token 时,如果全量重算,这一层就要做 O(n²·d) 的矩阵乘法。

整个生成过程从 1 到 n,总计算量是 O(n³·d) 量级(每步 O(n²·d),共 n 步)。做了 KV Cache 之后,每步只需要算新 token 的 Q、K、V,以及新 Q 和缓存 K 的注意力,单步是 O(n·d),总共 O(n²·d)。差距是 n 倍。n 等于 1000 的时候,就是 1000 倍的差距,这还没算上内存带宽和 kernel 启动的开销。

提示:这里说的“计算量”是理论 FLOPs。实际推理里,decode 阶段往往是内存带宽受限而不是算力受限,KV Cache 减少的不只是计算,更重要的是减少了对历史 K、V 的重复读写。

2.4 一个直观的类比

你可以把自回归生成想象成滚雪球。每滚一圈,雪球变大一点。如果不做缓存,相当于每滚一圈都要把整个雪球重新捏一遍;做了缓存,相当于只把新粘上的那层雪捏上去,原来的球体保持不动。雪球越大,省下的力气越多。KV Cache 就是这个“保持不动的球体”。

3. KV Cache 的核心原理与显存账本

3.1 缓存的结构:每层、每头、每位置

KV Cache 不是一份,而是每一层都有一份。因为 Transformer 每一层的注意力参数不同,算出来的 K、V 也不同,不能跨层复用。

在每一层内部,如果是多头注意力(Multi-Head Attention),每个头也有自己独立的 K、V。所以缓存的形状大致是:

[num_layers, 2, batch_size, num_heads, seq_len, head_dim]

其中那个 2 就是 K 和 V。有些实现会把 K 和 V 分开存成两个张量,有些会拼在一起,逻辑上一样。

3.2 显存占用怎么估

这是工程上最关心的问题。单个 token 的 KV Cache 大小可以这样算:

每 token 字节数 = 2 × num_layers × num_heads × head_dim × dtype_bytes

注意num_heads × head_dim通常等于隐藏维度 d(在标准多头里)。所以也可以写成:

每 token 字节数 = 2 × num_layers × d × dtype_bytes

举个具体例子。假设一个模型 L=32 层,隐藏维度 d=4096,用 FP16(2 字节)存储:

每 token = 2 × 32 × 4096 × 2 = 524288 字节 ≈ 0.5 MB

也就是说,每生成一个 token,KV Cache 就要多占 0.5 MB。如果上下文长度到 8192,batch size 为 1,那就是:

0.5 MB × 8192 ≈ 4 GB

这还只是 batch=1。如果并发 16 路,直接 64 GB 显存没了。这就是为什么长上下文 + 高并发是显存杀手。

参数符号示例值
层数L32
隐藏维度d4096
数据类型dtypeFP16 (2B)
每 token 缓存-0.5 MB
序列长度n8192
batch sizeb1
总缓存-约 4 GB

3.3 prefill 和 decode 两个阶段的差异

KV Cache 的引入,把推理明确分成了两个阶段:

Prefill 阶段:处理用户输入的 prompt。这时候所有 token 都是已知的,可以并行计算,一次性把整段 prompt 的 K、V 算出来填进缓存。这个阶段是计算密集型,GPU 利用率高。

Decode 阶段:逐个生成新 token。每步只算一个新 token 的 Q、K、V,然后拿新 Q 去和缓存里所有 K 做注意力。这个阶段是内存带宽密集型,因为每步都要把整个 KV Cache 读一遍。

这两个阶段的性能特征完全不同,优化手段也不一样。很多推理框架会把它们分开调度,甚至用不同的 kernel。理解这一点,对排查“为什么首 token 慢”和“为什么后续吐字慢”很有帮助。

3.4 为什么只缓存 K 和 V,不缓存 Q

前面提过,Q 是“当前 token 用来查询的向量”,每个位置只在它自己被生成的那一步用一次,之后再也不用了。缓存 Q 没有任何复用价值,只会白白占显存。而 K 和 V 是“被查询的对象”,会被后续所有 token 反复读取,所以必须缓存。

这个设计不是随便定的,是从注意力计算的数学结构里推出来的必然结果。你只要记住一句话:Q 是一次性的,K 和 V 是长期被引用的。

4. 动手实现:从零写一个带 KV Cache 的注意力

4.1 不带缓存的朴素实现

先看一个最朴素的单步注意力,方便对比:

import torch import torch.nn.functional as F def attention_no_cache(x, W_q, W_k, W_v, past_kv=None): # x: [batch, seq_len, d] Q = x @ W_q K = x @ W_k V = x @ W_v d_k = Q.size(-1) scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) attn = F.softmax(scores, dim=-1) out = attn @ V return out

每次调用都要把完整序列传进来,K、V 全量重算。生成 100 个 token 就调用 100 次,每次序列都在变长。

4.2 带 KV Cache 的实现

改造思路很简单:把历史 K、V 存起来,每次只算新 token 的 Q、K、V,然后把新 K、V 拼到缓存后面。

def attention_with_cache(x_new, W_q, W_k, W_v, past_kv=None): # x_new: [batch, 1, d] 只包含当前新 token Q = x_new @ W_q K_new = x_new @ W_k V_new = x_new @ W_v if past_kv is not None: K_past, V_past = past_kv K = torch.cat([K_past, K_new], dim=1) V = torch.cat([V_past, V_new], dim=1) else: K, V = K_new, V_new d_k = Q.size(-1) scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) attn = F.softmax(scores, dim=-1) out = attn @ V new_kv = (K, V) return out, new_kv

关键点有三个:一是输入只传新 token,二是用torch.cat把新 K、V 接到历史后面,三是把更新后的缓存返回出去,供下一步使用。

4.3 多头版本要注意的维度

多头注意力的缓存形状是[batch, num_heads, seq_len, head_dim]。拼接的时候要拼在seq_len那一维,也就是dim=2,别拼错。很多新手第一次写会拼到 head 维度上,结果注意力全乱。

# K_past: [batch, num_heads, past_len, head_dim] # K_new: [batch, num_heads, 1, head_dim] K = torch.cat([K_past, K_new], dim=2)

4.4 预分配缓存 vs 动态拼接

上面用torch.cat是最直观的写法,但工程上很少这么干。因为cat每次都会新分配一块内存并拷贝,序列长了之后开销很大,还会造成显存碎片。

主流做法是预分配一块固定大小的缓存,形状是[batch, num_heads, max_seq_len, head_dim],然后用一个位置指针记录当前写到哪了。每步只往对应位置写,不重新分配。

# 预分配 cache_k = torch.empty(batch, num_heads, max_len, head_dim, device=device) cache_v = torch.empty_like(cache_k) pos = 0 # 写入 cache_k[:, :, pos:pos+1, :] = K_new cache_v[:, :, pos:pos+1, :] = V_new pos += 1 # 读取时只取前 pos 个 K = cache_k[:, :, :pos, :] V = cache_v[:, :, :pos, :]

这个改动看起来小,但对吞吐的影响很大。预分配避免了反复分配释放,也让内存访问更连续。

注意:预分配的最大长度要提前定好。如果实际序列超过这个长度,要么截断,要么扩容。扩容时机的选择是个工程权衡,扩太早浪费显存,扩太晚触发重分配卡顿。

4.5 位置编码的配合

KV Cache 缓存的是 K、V 的数值,但位置信息是在算 K、V 之前就注入进去的(比如 RoPE 旋转位置编码)。所以缓存里的 K 已经带了位置信息,后续直接复用没问题。但如果你用的是绝对位置编码,且实现方式是“在注意力分数上加位置偏置”,那就要小心:新 token 的 Q 和缓存 K 做分数计算时,位置偏置要按各自的绝对位置来算,不能简单复用。

RoPE 之所以在长上下文模型里流行,一个原因就是它和 KV Cache 配合得很自然——位置信息编码在 K 里,缓存即用,不需要额外处理。

5. 工程实践中的坑与优化手段

5.1 显存不够怎么办:几个方向

显存是 KV Cache 最直接的约束。常见应对手段有这么几类:

减少并发:最粗暴,但吞吐直接掉。适合延迟敏感、吞吐不敏感的场景。

缩短上下文:限制 max_seq_len,或者做滑动窗口注意力,只保留最近 N 个 token 的缓存。代价是丢失远距离信息。

量化缓存:把 KV Cache 从 FP16 降到 INT8 甚至 INT4。显存直接减半或减到四分之一,精度损失通常可控,但需要校准。这是目前性价比很高的手段。

分组查询注意力(GQA):让多个 Query 头共享一组 K、V 头。比如 32 个 Q 头只配 8 个 KV 头,缓存直接降到四分之一。现在很多开源模型默认就用 GQA,就是为了省这块显存。

分页管理:借鉴操作系统的虚拟内存思路,把缓存切成固定大小的块(page),按需分配,减少碎片。这个思路在主流推理框架里已经很成熟。

手段显存收益主要代价
减少并发线性吞吐下降
滑动窗口与窗口成正比丢失长距离信息
量化缓存2x ~ 4x精度损失、需校准
GQA与头数比成正比表达能力略降
分页管理减少碎片实现复杂度

5.2 缓存复用:前缀共享

如果你的服务里有很多请求共享同一段前缀(比如相同的系统提示词),那这段前缀的 KV Cache 是可以复用的。第一个请求算完,把前缀部分的缓存留下来,后续请求直接接上,省掉重复的 prefill。

这个优化在多轮对话里特别有用:历史对话的缓存可以保留,新一轮只 prefill 新增的用户输入。实现上需要一个前缀匹配和缓存索引机制,复杂度不低,但收益很可观。

5.3 常见问题速查

现象可能原因排查方向
生成越来越慢缓存没生效,每步全量重算检查是否每步都传了完整序列
显存随长度暴涨缓存未预分配,碎片严重看显存分配曲线是否锯齿状
输出乱码/重复缓存拼接维度错检查 cat 的 dim 是否为 seq 维
首 token 慢后续快prefill 计算密集正常现象,优化 prefill kernel
并发上不去缓存占用过大估算每 token 字节数,考虑量化或 GQA
长文本质量下降滑动窗口截断了关键信息调整窗口大小或换注意力方案

5.4 几个容易忽略的细节

第一,缓存的清理时机。请求结束后要及时释放对应的缓存块,否则显存会慢慢泄漏。分页管理里通常有个引用计数,归零就回收。

第二,batch 内不同序列长度不一致。同一个 batch 里,有的请求已经生成 100 个 token,有的才 10 个。这时候缓存的有效长度不同,注意力计算要做 mask,否则短序列会读到别的序列的缓存。这个 bug 很隐蔽,输出可能看起来“差不多对”,但质量会悄悄下降。

第三,数值精度。缓存长时间累积,如果中间有精度损失,误差会逐步放大。FP16 缓存在超长序列下可能出现数值不稳定,有些场景会保留一份 FP32 的累加。这个要看具体模型和任务,不是所有场景都需要。

第四,beam search 下的缓存管理。beam search 会同时维护多条候选路径,每条路径有自己的缓存。beam 合并、剪枝的时候,缓存也要跟着合并和丢弃,逻辑比贪心解码复杂不少。如果实现不当,很容易出现缓存和路径对不上的问题。

6. 从 KV Cache 延伸出去的几个思考

KV Cache 看起来只是“存一下 K 和 V”,但它其实是整个大模型推理优化的一个缩影。它把“计算换存储”这个经典权衡摆在了台面上:不做缓存,算力扛不住;做了缓存,显存扛不住。所有的优化手段,本质上都是在找这两者之间的平衡点。

我自己的体会是,理解 KV Cache 最好的方式不是背公式,而是亲手写一遍带缓存的注意力,然后跑一个长序列生成,看着显存曲线和延迟曲线,你自然就明白每一步在发生什么。纸上推一百遍 O(n²) 和 O(n),不如实际 profile 一次。

另外,KV Cache 的设计也影响模型架构的选择。为什么现在新模型越来越多用 GQA、MQA,为什么 RoPE 成了标配,为什么长上下文模型要专门设计注意力模式——这些决策背后都有 KV Cache 的影子。从这个角度说,搞懂 KV Cache,不只是搞懂一个优化技巧,而是搞懂了现代大模型推理的一条主线。

如果你正在调推理服务,建议先把每 token 的缓存字节数算清楚,再对照你的显存和并发目标,看看缺口有多大。这个数字一出来,该量化还是该换注意力方案,方向基本就定了。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/27 5:42:14

从0到8192:基于期望搜索与启发式评估的2048游戏AI实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/27 5:37:41

引用抽检别只会扫一眼:同一条文献可勾选四项核验骨架

参考文献最常见的假完成是扫一眼:列表格式整齐、作者年份齐全、期刊名看着眼熟,于是整页放行。可答辩或外审真正追问的往往是另外几件事:这条能不能当场找到原文?列表写的年份和原文是否一致?正文那句「已有研究表明」…

作者头像 李华
网站建设 2026/9/27 5:31:27

用I2C获取的数据经常卡死的原因

1.I2C时序容易被程序中的中断给打断,并且程序中多个任务同时运行,对时序要求高,I2C数据就可能卡死时序的概念:通信的时候,什么时间发信号,什么时间发数据,谁先谁后,这一整套先后顺序…

作者头像 李华
网站建设 2026/9/27 5:30:21

Jetson Orin NX稳定接入联适R70M-GNSS串口调优全指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/27 5:29:29

嵌入式C++开发工具链解剖:从CubeMX到GCC再到Keil与VS Code的协同本质

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华