自回归生成的隐含浪费
大模型生成文本是自回归的:喂入提示,算出下一个 token;把这个 token 接到末尾,再算下一个;循环往复。每一步只新增一个 token,但注意力机制要求新 token 和它之前的所有 token 都交互一遍——历史越长,每一步要「回看」的就越多。
这个「回看」如果每次都从头算,代价会失控。注意力需要每个历史 token 的两组中间向量:K(key,被检索的键)和 V(value,被取回的内容),它们由每层的权重投影得到。没有缓存时,生成第 t 个 token 就要把前面 t-1 个 token 的 K 和 V 重新投影一遍——这些向量上一步刚算过,一步之后原封不动地再算一遍。把整段生成加起来,重复计算随长度平方级膨胀:生成 n 个 token,历史 K/V 的重算量正比于 n²,n 到千级时就是数十万次 token 级的重算,其中绝大部分是纯浪费。
KV Cache:用显存换计算
解法直观:K 和 V 只依赖 token 本身与层权重,token 不变它们就不变,那就算一次、存下来。把每一层、每个历史 token 的 K 和 V 缓存在显存里,就是 KV Cache。有了它,每步生成只做三件事:算新 token 自己的 q、k、v 三个投影;把 k、v 追加进缓存;用新 q 对整个缓存做一次注意力。历史 token 的所有重算消失,整段生成的总计算量从平方级降到线性。
顺带一提,即使有缓存,注意力每步的计算量仍随历史长度线性增长:回看的对象越来越多,单步延迟随之上升。缓存省掉的是重复计算,省不掉「看得越来越多」这件事本身。
这是典型的用显存换计算,而账单随长度增长。
显存账:一个值得背下来的公式
KV Cache 的显存估算公式只有一行:
每 token 字节数 = 2 × 层数 × KV 头数 × 每头维度 × 每元素字节数
公式里那个 2 是 K 和 V 两组缓存;「每元素字节数」在 fp16 下取 2。头数与每头维度由模型结构决定——多头注意力(MHA)里 KV 头数等于查询头数,而 GQA 模型的 KV 头数少得多,这正是它省显存的地方。
代入一个典型 7B 配置做估算:32 层、32 个 KV 头、每头 128 维、fp16。每 token = 2 × 32 × 32 × 128 × 2 = 512 KB。于是单个序列:1K 上下文约 0.5 GB,4K 约 2 GB,32K 约 16 GB,128K 约 64 GB。对比一下:7B 模型的权重用 fp16 存大约 14 GB。也就是说,长序列下一条请求的 KV Cache 就能超过模型权重本身。以上均为量级估算,实际数字随模型结构、精度和批大小变化——批大小是乘法因子,八路并发就是八倍。公式给出的还是下界:某些实现要额外暂存注意力分数矩阵,实际占用只会更高。
为什么长上下文是推理瓶颈
有了这笔账,很多工程现象立刻说得通。长上下文服务贵,贵在 KV Cache 的显存占用和随之下降的单卡并发数:显存被缓存吃掉,能同时服务的请求数就少。预填充(prefill,用户输入的整段一次性进模型算好缓存)和逐 token 解码(decode)之所以被分开讨论,也因为两者对算力与缓存的压力分布完全不同。上下文窗口竞赛的成本壁垒,很大程度上就是 KV Cache 的壁垒。服务商推出的「上下文缓存」(相同前缀只算一次、后续请求复用),本质就是把 KV Cache 从请求内复用扩大到请求间复用——重复的系统提示、长文档不再每次重算,费用与延迟都按复用打折。理解了缓存的意义,这类定价机制就不神秘了。
三个优化方向
第一,GQA/MQA:让多个查询头共享同一组 KV 头(MQA 是全体共享一组,GQA 是分组共享),KV Cache 成比例缩小,质量损失在多数场景可接受,已是主流模型的标配。第二,KV 量化:把缓存从 fp16 压到 8 bit 甚至更低精度,显存直接减半再减半,代价是长序列下的精度敏感性。第三,PagedAttention:借鉴操作系统虚拟内存的分页思想,把每个请求的缓存切成小块非连续存放,消除按最大长度预分配的浪费,vLLM 等推理引擎靠它把并发吞吐提升了两到四倍(论文口径,相对同期系统)。
这三个方向分别动公式里的三个变量:KV 头数、每元素字节数、序列长的管理方式。每一项都值得单独写一篇,这里留作钩子。
小结
自回归生成天然带着「每步重算全部历史」的浪费,KV Cache 用一次投影、永久缓存把它抹平,代价是随长度线性增长的显存。记住那个公式和 512 KB 这个量级,你就有了一把判断推理优化的尺子:任何新方案,先问它优化的是公式里的哪一项。
读者留言
COMMENTS 暂无还没有留言,来说第一句?