1. 从一个让人抓狂的推理速度问题说起如果你自己动手跑过自回归生成模型比如用 GPT 系列或者 LLaMA 系列做文本续写大概率遇到过这样一个现象生成第一个 token 的时候速度还行越往后越慢生成到几百个 token 之后每个 token 的耗时几乎是线性增长的。我最早在单卡上跑一个 7B 模型做对话测试前 50 个 token 每秒能出十几个到第 500 个 token 的时候直接掉到每秒两三个体验非常糟糕。这个问题困扰了我挺长时间直到我把注意力计算的复杂度认认真真推导了一遍才意识到瓶颈根本不在模型本身而在于每一步都在重复计算已经算过的东西。KV CacheKey-Value Cache就是解决这个问题的核心手段它几乎是所有现代自回归大模型推理框架的标配。这篇文章我会从注意力机制的底层计算出发把 KV Cache 是什么、为什么需要它、怎么实现它、以及实际用的时候有哪些坑全部拆开讲清楚。不管你是刚接触 Transformer 的新手还是已经在做推理优化的工程师应该都能从中拿到一些能直接用的东西。先给一个最直白的定义KV Cache 是在自回归生成过程中把每一层注意力计算里已经产生过的 Key 和 Value 张量缓存下来后续步骤直接复用避免重复计算。它不改变模型结构不改变输出结果纯粹是一个用显存换时间的工程优化。理解它的前提是先把自注意力机制里 Q、K、V 三个矩阵的角色和计算流程搞清楚。2. 自注意力机制回顾与 KV Cache 的切入点2.1 Q、K、V 到底在算什么Transformer 的自注意力机制本质上是一个查询-匹配-聚合的过程。输入序列经过三个线性变换分别得到 Query查询、Key键、Value值三个矩阵。用生活化的类比来说Query 就像你在搜索引擎里输入的关键词Key 像每个网页的标题和标签Value 像网页的实际内容。注意力计算做的事情就是拿你的查询去和所有网页的标签做匹配算出匹配度然后按匹配度加权汇总网页内容。具体到公式层面缩放点积注意力的计算是Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) V其中d_k是 Key 向量的维度除以sqrt(d_k)是为了防止点积结果过大导致 softmax 梯度消失。这个公式里Q K^T得到的是一个形状为(seq_len_q, seq_len_k)的注意力分数矩阵经过 softmax 归一化后再和 V 相乘得到最终的输出。在自回归生成场景下模型是一步一步往外吐 token 的。假设当前已经生成了 t 个 token现在要生成第 t1 个 token。这时候 Query 只有最后一个位置的那一个向量因为只有新 token 需要查询但 Key 和 Value 需要覆盖从第 1 个到第 t 个所有位置。问题就出在这里第 1 到第 t 个位置的 Key 和 Value在之前的每一步里其实都已经算过了。2.2 没有 Cache 时到底浪费了多少计算我们来算一笔账。假设序列长度为 n隐藏维度为 d注意力头数为 h每个头的维度为 d_head d/h。在生成第 i 个 token 时如果不做任何缓存需要重新计算前 i 个位置的所有 K 和 V。单层单头的情况下计算 K 和 V 的矩阵乘法量级大约是O(i * d * d_head)。把 i 从 1 累加到 n总计算量是O(n^2 * d * d_head)。而如果做了缓存每一步只需要计算新 token 的 K 和 V单步计算量是O(d * d_head)累加起来是O(n * d * d_head)。也就是说不做缓存的情况下K/V 的计算量是缓存情况下的 n 倍。当 n 等于 1000 的时候这就是 1000 倍的差距虽然实际中还有其他计算开销摊薄这个比例但量级上的差异是实打实的。我第一次意识到这个问题是在 profiling 一个推理脚本的时候发现k_proj和v_proj这两个线性层的调用次数随着生成步数线性增长而q_proj始终只调用一次。当时就觉得不对劲后来才明白这就是没有缓存的典型表现。加上缓存之后k_proj和v_proj每步也只调用一次整体耗时曲线立刻从线性增长变成了基本平坦。2.3 KV Cache 的核心思想一句话说清KV Cache 的核心思想可以用一句话概括把历史 token 的 Key 和 Value 存起来每步只计算新 token 的 Key 和 Value然后拼接到缓存后面。这样注意力的计算就变成了新 Query 查询全部历史 Key而不是全部 Query 查询全部 Key。这里有一个关键点需要强调Query 是不需要缓存的。因为在自回归生成中每一步只有最新位置的 token 需要输出历史位置的 Query 已经完成了它们的使命不会再被用到。只有 Key 和 Value 会被后续步骤反复引用所以只需要缓存这两个。这也是为什么它叫 KV Cache 而不是 QKV Cache。3. KV Cache 的具体实现与显存开销分析3.1 缓存的形状与拼接逻辑在实际代码里KV Cache 通常是一个形状为(batch_size, num_heads, max_seq_len, head_dim)的张量每一层注意力都有自己的缓存。以 HuggingFace 的transformers库为例past_key_values就是一个嵌套结构外层是层数内层是 key 和 value 两个张量。拼接逻辑大概是这样的每一步生成时先计算当前 token 的 k 和 v形状是(batch_size, num_heads, 1, head_dim)然后沿着序列维度第 2 维拼接到已有的缓存上得到(batch_size, num_heads, cur_len, head_dim)。接着用当前 token 的 q形状(batch_size, num_heads, 1, head_dim)去和拼接后的 k 做点积得到注意力分数再和 v 加权求和。用 PyTorch 伪代码表示大概是# 假设 past_k, past_v 是上一步的缓存形状 (B, H, L, D) # 当前步计算得到新的 k, v形状 (B, H, 1, D) k torch.cat([past_k, k_new], dim2) # (B, H, L1, D) v torch.cat([past_v, v_new], dim2) # (B, H, L1, D) # 注意力计算q 形状 (B, H, 1, D) attn torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(head_dim) attn torch.softmax(attn, dim-1) out torch.matmul(attn, v) # (B, H, 1, D)注意这里的torch.cat操作它会在每一步都创建一个新的张量带来额外的内存分配和拷贝开销。在生产级推理框架里通常会预分配一块足够大的显存然后用索引写入的方式更新缓存避免频繁的 cat 操作。这一点后面讲优化的时候会展开。3.2 显存占用到底有多大KV Cache 的显存占用公式是显存占用 2 * batch_size * num_layers * num_heads * max_seq_len * head_dim * dtype_size前面的 2 是因为要同时存 Key 和 Value。我们拿一个具体的模型来算。假设是 LLaMA-2 7B配置是 32 层32 个注意力头head_dim 为 128隐藏维度 4096。用 FP16 存储每个元素 2 字节。单条序列、长度为 2048 的情况下2 * 1 * 32 * 32 * 2048 * 128 * 2 字节 2 * 32 * 32 * 2048 * 128 * 2 2 * 32 * 32 * 2048 * 256 2 * 32 * 32 * 524288 2 * 32 * 16777216 1073741824 字节 ≈ 1 GB也就是说单条 2048 长度的序列KV Cache 就要吃掉大约 1GB 显存。如果 batch_size 是 8那就是 8GB。而模型本身的权重7B 参数FP16大约是 14GB。加起来 22GB一张 24GB 的卡基本就满了。这就是为什么长上下文场景下显存这么紧张。我把不同配置下的显存占用整理成一张表方便你对照模型规模层数头数head_dim序列长度batch_sizeKV Cache 显存7B323212820481约 1.0 GB7B323212840961约 2.0 GB7B323212881921约 4.0 GB13B404012840961约 3.1 GB70B806412840961约 10.0 GB可以看到序列长度翻倍KV Cache 显存就翻倍这是线性增长。而 batch_size 翻倍显存也翻倍。所以在实际部署中KV Cache 的显存管理直接决定了你能支持多长的上下文、多大的并发。3.3 为什么不用缓存 Query前面提到 Query 不需要缓存这里再补充一下原因。在自回归生成中每一步的输出只依赖于当前位置的 Query 和所有历史位置的 Key/Value。历史位置的 Query 对应的输出已经在之前的步骤中产生过了不会再被使用。缓存 Query 不仅没有收益还会白白占用显存。所以 KV Cache 只缓存 K 和 V这是经过严格推导的最优选择。4. 从零实现一个带 KV Cache 的注意力层4.1 基础版本实现光看公式不够直观我直接写一个可以跑的最小实现。下面这段代码实现了一个带 KV Cache 的多头注意力层用 PyTorch 写的去掉了所有不必要的封装方便你看清楚每一步在做什么。import torch import torch.nn as nn import math class CachedAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.d_model d_model self.num_heads num_heads self.head_dim d_model // num_heads assert self.head_dim * num_heads d_model self.q_proj nn.Linear(d_model, d_model) self.k_proj nn.Linear(d_model, d_model) self.v_proj nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) def forward(self, x, past_kvNone): # x: (B, T, D)T 是当前步的序列长度 B, T, D x.shape q self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) k self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) v self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) # q, k, v: (B, H, T, head_dim) if past_kv is not None: past_k, past_v past_kv k torch.cat([past_k, k], dim2) v torch.cat([past_v, v], dim2) # 保存当前缓存 new_kv (k, v) # 注意力计算 attn torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) # 因果掩码防止看到未来 if T 1: mask torch.triu(torch.ones(T, k.size(2)), diagonal1).bool().to(x.device) attn attn.masked_fill(mask, float(-inf)) attn torch.softmax(attn, dim-1) out torch.matmul(attn, v) # (B, H, T, head_dim) out out.transpose(1, 2).contiguous().view(B, T, D) return self.out_proj(out), new_kv这段代码里有两个细节值得注意。第一past_kv是作为参数传入的而不是存在模块内部这样设计是为了让缓存的生命周期由调用方控制避免状态污染。第二因果掩码只在 T 大于 1 的时候才需要因为生成阶段 T 通常等于 1此时不需要掩码。4.2 预分配缓存优化版本上面的基础版本每一步都在做torch.cat会频繁分配新内存。在生产环境里更好的做法是预分配一块固定大小的缓存然后用索引写入。下面是一个优化版本的示意class PreallocatedCachedAttention(nn.Module): def __init__(self, d_model, num_heads, max_seq_len): super().__init__() self.d_model d_model self.num_heads num_heads self.head_dim d_model // num_heads self.max_seq_len max_seq_len self.q_proj nn.Linear(d_model, d_model) self.k_proj nn.Linear(d_model, d_model) self.v_proj nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) def init_cache(self, batch_size, device, dtype): # 预分配缓存形状 (B, H, max_seq_len, head_dim) k_cache torch.zeros( batch_size, self.num_heads, self.max_seq_len, self.head_dim, devicedevice, dtypedtype ) v_cache torch.zeros_like(k_cache) return k_cache, v_cache def forward(self, x, cache, start_pos): B, T, D x.shape q self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) k self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) v self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) k_cache, v_cache cache # 直接写入预分配的位置避免 cat k_cache[:, :, start_pos:start_pos T, :] k v_cache[:, :, start_pos:start_pos T, :] v # 只取有效部分参与计算 k_valid k_cache[:, :, :start_pos T, :] v_valid v_cache[:, :, :start_pos T, :] attn torch.matmul(q, k_valid.transpose(-2, -1)) / math.sqrt(self.head_dim) if T 1: mask torch.triu( torch.ones(T, start_pos T), diagonalstart_pos 1 ).bool().to(x.device) attn attn.masked_fill(mask, float(-inf)) attn torch.softmax(attn, dim-1) out torch.matmul(attn, v_valid) out out.transpose(1, 2).contiguous().view(B, T, D) return self.out_proj(out)这个版本的关键改进是k_cache[:, :, start_pos:start_pos T, :] k这一行它直接在预分配的显存上做原地写入没有额外的内存分配。实测下来在长序列生成场景下这个改动能带来 15% 到 30% 的吞吐提升具体取决于序列长度和硬件。4.3 缓存管理的几个实操要点实现 KV Cache 的时候有几个坑我踩过这里直接列出来缓存的 dtype 要和模型一致。如果模型是 FP16缓存也应该是 FP16混用 FP32 会让显存翻倍而且可能触发类型转换的额外开销。注意缓存的设备一致性。多卡推理时缓存必须和对应的层在同一张卡上跨卡拷贝的代价极高。序列结束时要重置缓存。如果一批请求处理完不清理缓存下一批请求会读到脏数据输出会完全错乱。我见过有人因为这个 bug 排查了一整天。预分配的 max_seq_len 要留余量。如果实际序列超过预分配长度写入会越界报错。通常建议按最大可能长度的 1.2 倍来分配。提示调试 KV Cache 相关问题时可以先关闭缓存跑一遍再打开缓存跑一遍对比输出是否完全一致。如果一致说明缓存逻辑正确如果不一致大概率是拼接位置或者掩码出了问题。5. 进阶优化与常见问题排查5.1 多查询注意力与分组查询注意力标准的多头注意力里每个头都有独立的 K 和 V缓存大小和头数成正比。为了压缩 KV Cache研究者提出了两种变体多查询注意力MQA和分组查询注意力GQA。MQA 的做法是所有头共享同一组 K 和 V只有 Q 保持多头。这样 KV Cache 直接缩小到原来的 1/num_heads。GQA 是折中方案把头分成若干组每组共享一组 K 和 V。比如 32 个头分成 8 组每组 4 个头共享缓存就缩小到 1/4。这两种方案在 LLaMA-2 70B、Mistral 等模型里都有应用。实测下来GQA 在几乎不损失效果的前提下能把 KV Cache 显存降低 4 到 8 倍对长上下文场景非常友好。如果你在选型阶段优先考虑带 GQA 的模型。5.2 分页缓存与 PagedAttention即使有了 GQA长上下文下的显存碎片问题依然存在。传统做法是给每条序列预分配一块连续显存但序列实际长度参差不齐预分配会导致大量浪费。vLLM 提出的PagedAttention借鉴了操作系统的虚拟内存分页思想把 KV Cache 切成固定大小的块block按需分配用块表来管理映射关系。这个方案的好处是显存利用率能从 20% 到 40% 提升到 90% 以上同时支持不同序列共享块比如相同的 prompt 前缀。如果你在做高并发推理服务PagedAttention 几乎是必选项。不过它的实现复杂度较高自己从零写不太现实建议直接用 vLLM 或者 TensorRT-LLM 这类成熟框架。5.3 常见问题速查表下面这张表整理了我在实际使用 KV Cache 过程中遇到的高频问题以及对应的排查思路问题现象可能原因排查方法解决方案输出结果和不开缓存不一致缓存拼接位置错误对比每步的 k/v 张量检查 cat 的 dim 参数显存溢出缓存未释放或预分配过大打印缓存形状和显存占用减小 max_seq_len 或 batch_size生成速度没有提升缓存未生效profiling 看 k_proj 调用次数确认 past_key_values 正确传递长序列后输出乱码缓存越界或 dtype 不匹配检查序列长度是否超限增大预分配长度统一 dtype多卡推理结果异常缓存设备不一致打印缓存所在 device确保缓存和层在同一设备批处理时部分序列出错缓存未按序列隔离检查 batch 维度处理每条序列独立管理缓存5.4 几个容易被忽略的细节第一个细节是位置编码和缓存的配合。用旋转位置编码RoPE的模型缓存的 K 在写入时已经应用了位置编码后续直接复用即可不需要重新编码。但如果用的是可学习的位置嵌入就要注意缓存里的 K 是否包含了正确的位置信息。这一点在自定义实现时特别容易出错。第二个细节是缓存的量化。为了进一步压缩显存可以把 KV Cache 量化到 INT8 甚至 INT4。实测下来INT8 量化对效果的影响很小显存直接减半。但量化会引入反量化开销需要权衡。如果显存是瓶颈量化是值得的如果算力是瓶颈可能得不偿失。第三个细节是前缀共享。在多轮对话场景下系统提示词system prompt往往是固定的。如果能把这段前缀的 KV Cache 缓存下来所有请求共享能省下大量重复计算。vLLM 的 prefix caching 就是做这个的。我自己实现过一个简化版把 system prompt 的缓存存成文件启动时加载实测首 token 延迟降低了 40% 以上。6. 一些实战中的经验体会KV Cache 这个东西看原理的时候觉得很简单不就是把 K 和 V 存起来嘛。但真正在工程里用好需要考虑的东西远比想象中多。显存管理、批处理调度、缓存生命周期、量化精度、多卡同步每一个环节都有坑。我个人的经验是如果你只是做实验或者小规模推理直接用 HuggingFace 的generate接口它内部已经帮你处理好了 KV Cache你只需要传use_cacheTrue就行。但如果你要做生产级部署或者要针对特定场景做极致优化那就必须深入理解缓存的每一个细节甚至自己动手改底层实现。还有一个体会是不要过早优化。我见过有人一上来就搞 PagedAttention、量化缓存结果基础逻辑都没跑通排查问题的时候根本不知道是哪一层出的错。正确的做法是先跑通基础版本确认输出正确再逐步加优化每加一层都做对比测试。这样出问题的时候范围能缩小到最近一次改动。最后分享一个调试小技巧在缓存的写入和读取位置各加一个断言检查形状和数值范围。比如写入后断言k_cache[:, :, start_pos:start_posT, :].abs().max() 100读取时断言k_valid.shape[2] start_pos T。这些断言在开发阶段能帮你快速定位问题上线前再去掉。我自己靠这个习惯省下了大量排查时间。
企业数字化 ERP 产品动态
相关推荐
5G核心网SBA与CUPS实战解析:从协议到现网故障定位 /* 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 2:24:28
STM32H750 ADC+DMA+定时器实现高精度电压监测实战 /* 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 2:24:28
自己改论文后AI率反而升高,怎样修改才能降低AI率? 自己改论文后AI率反而升高,怎样修改才能降低AI率?
你没有让AI替写,而是自己删套话、换句子、合并段落。改完觉得比原来简洁,AI率却上升了,连没有动过的地方似乎也出现新提示。继续自己改担心再涨,恢复原文… · 2026/9/27 2:24:22
ARM开发板安装ROS2 Humble完整指南:从换源到避坑实战 /* 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 2:59:40
智能体通信协议MCP/A2A/ANP选型与落地实践指南 /* 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 2:59:34
STM32项目源码实战:从环境搭建到避坑调试全指南 /* 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 2:59:27
Chrome网课视频自动暂停原因与防暂停扩展解决方案 /* 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 2:59:09
监控项目等不起!物联网卡当天发货、分钟级监控,工程商该知道 做安防工程的朋友这两年应该都有同一个感受:项目变少了,利润变薄了,甲方对交付速度和后期运维的要求反而越来越高。以前装完摄像头就算完事,现在甲方要看到画面流畅、数据稳定、出了问题有人快速响应。项目本身已经不赚钱了&#… · 2026/9/27 2:59:09
长春网站建设网站源码怎么改才不卡:性能优化实战指南 长春网站建设网站源码怎么改才不卡:性能优化实战指南 别再说模板网站太丑了,更可怕的是打开要等5秒,客户直接关掉。 很多长春本地企业老板找我们要 长春网站建设网站源码 ,核心诉求就一个:别卡顿,要快。 其实源码本身不慢,慢的是你没懂… · 2026/9/27 2:59:09
MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现 简介:这套Matlab仿真工具完整呈现雷达信号脉冲压缩过程,从线性调频(LFM)信号生成、目标回波仿真到匹配滤波压缩处理均有可运行代码支撑,面向电子信息工程、计算机、数学等专业学生,适用于课程设计、期末大作… · 2026/9/27 0:00:01
汕头网站建设制作厂家避坑指南:5大注意事项救急 汕头网站建设制作厂家避坑指南:5大注意事项救急 改个需求建站公司拖一周,这种憋屈事我见得太多了。 很多汕头老板找本地建站团队,签合同前看着方案挺美,一上线就变脸。 今天不聊虚的,直接拆解找 汕头网站建设制作厂家 时的5个核心 注意事项… · 2026/9/27 0:00:01
多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习 简介:基于PyTorch的多模态虚假新闻检测项目完整代码包,面向自然语言处理与计算机视觉交叉方向的开发者、科研人员及毕业设计选题者,解决社交媒体中文本与图像联合识别虚假新闻的问题。系统以BERT预训练模型提取文本语义特征,以Res… · 2026/9/27 0:00:01
MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现 简介:这套Matlab仿真工具完整呈现雷达信号脉冲压缩过程,从线性调频(LFM)信号生成、目标回波仿真到匹配滤波压缩处理均有可运行代码支撑,面向电子信息工程、计算机、数学等专业学生,适用于课程设计、期末大作… · 2026/9/27 0:00:01
汕头网站建设制作厂家避坑指南:5大注意事项救急 汕头网站建设制作厂家避坑指南:5大注意事项救急 改个需求建站公司拖一周,这种憋屈事我见得太多了。 很多汕头老板找本地建站团队,签合同前看着方案挺美,一上线就变脸。 今天不聊虚的,直接拆解找 汕头网站建设制作厂家 时的5个核心 注意事项… · 2026/9/27 0:00:01
多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习 简介:基于PyTorch的多模态虚假新闻检测项目完整代码包,面向自然语言处理与计算机视觉交叉方向的开发者、科研人员及毕业设计选题者,解决社交媒体中文本与图像联合识别虚假新闻的问题。系统以BERT预训练模型提取文本语义特征,以Res… · 2026/9/27 0:00:01