拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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

1. 从一次推理延迟排查说起KV Cache 到底是什么前阵子帮朋友看一个文本生成的推理服务现象很典型单条请求响应挺快一旦并发上来延迟直接飙到没法用GPU 显存还涨得厉害。把 profiling 打开一看绝大部分时间耗在注意力计算上而且每生成一个 token前面所有 token 的键和值都被重新算了一遍。问题就出在这里——KV CacheKey-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、ValueQ 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²)注意力矩阵 QK^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举个具体例子。假设一个模型 L32 层隐藏维度 d4096用 FP162 字节存储每 token 2 × 32 × 4096 × 2 524288 字节 ≈ 0.5 MB也就是说每生成一个 tokenKV Cache 就要多占 0.5 MB。如果上下文长度到 8192batch size 为 1那就是0.5 MB × 8192 ≈ 4 GB这还只是 batch1。如果并发 16 路直接 64 GB 显存没了。这就是为什么长上下文 高并发是显存杀手。参数符号示例值层数L32隐藏维度d4096数据类型dtypeFP16 (2B)每 token 缓存-0.5 MB序列长度n8192batch sizeb1总缓存-约 4 GB3.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_kvNone): # 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_kvNone): # 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], dim1) V torch.cat([V_past, V_new], dim1) 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那一维也就是dim2别拼错。很多新手第一次写会拼到 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], dim2)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, devicedevice) cache_v torch.empty_like(cache_k) pos 0 # 写入 cache_k[:, :, pos:pos1, :] K_new cache_v[:, :, pos:pos1, :] 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 的缓存字节数算清楚再对照你的显存和并发目标看看缺口有多大。这个数字一出来该量化还是该换注意力方案方向基本就定了。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门