2.2 KV Cache 的显存账


文档摘要

2.2 KV Cache 的显存账 本节摘要:KV Cache 把算力预付成了显存。要算清这笔账,公式只有一项乘法:2 × 层数 × KV 头数 × 头维度 × 精度字节 × 序列长 × 批大小。本节逐因子拆开这个公式的来历,再用合理超参算出 7B 与 70B 在 4K/32K/128K 下的占用,最后点出一个反直觉结论——上下文足够长时,KV Cache 会比模型权重更先吃光显存。读懂这一节,你才算真正握住了"缓存是回收"的尺子。 显存到底被谁吃掉了 上一节我们把 KV Cache 想成"半成品仓库",没提仓库租金。现在该付账了。缓存要在显存里为每个已生成的 token 保留一份 K 矩阵和一份 V 矩阵,这两份矩阵永远不被释放,直到这个请求结束。

2.2 KV Cache 的显存账

本节摘要:KV Cache 把算力预付成了显存。要算清这笔账,公式只有一项乘法:2 × 层数 × KV 头数 × 头维度 × 精度字节 × 序列长 × 批大小。本节逐因子拆开这个公式的来历,再用合理超参算出 7B 与 70B 在 4K/32K/128K 下的占用,最后点出一个反直觉结论——上下文足够长时,KV Cache 会比模型权重更先吃光显存。读懂这一节,你才算真正握住了"缓存是回收"的尺子。

显存到底被谁吃掉了

上一节我们把 KV Cache 想成"半成品仓库",没提仓库租金。现在该付账了。缓存要在显存里为每个已生成的 token 保留一份 K 矩阵和一份 V 矩阵,这两份矩阵永远不被释放,直到这个请求结束。显存占用的总量,就是"一个 token 占多少"乘以"一共缓存了多少 token"。

一个 token 占多少,又由模型结构决定:每一层 Transformer 都有自己的一组注意力头,每个头投影出维度为"头维度"的 K 和 V 向量;所有层、所有头加起来,就是一个 token 在缓存里的全部体量。再乘上序列长度(缓存了多少 token)和批大小(同时缓存了多少个请求),就是整笔账。

关键特征先说在前面:权重是"一次买断、大小固定"的,买多大模型就占多少显存;KV Cache 是"随用随涨、按 token 计费"的,上下文越长、并发越高,它涨得越凶,而且没有上限封顶。这正是后面所有显存管理技巧的源头——它不是想优化,是被逼出来的。

把公式拆成六个因子

完整公式是:

KV 显存 = 2 × 层数 × KV头数 × 头维度 × 精度字节 × 序列长 × 批大小

逐个因子交代来历,免得死记:

  • 2:因为每个 token 要存两套矩阵,K 一套、V 一套。Q 不存,因为查询只在当前步算、用完即弃,不需要跨步复用。
  • 层数:每一层有独立的注意力头和独立的投影权重,因此每一层都要单独留一份 K、V。层数翻倍,缓存翻倍。
  • KV 头数:这里说的是"KV 头"而不是"注意力头总数"。在标准的多头注意力里两者相等;但分组查询注意力(GQA)和多人查询注意力(MQA)会让 KV 头远少于查询头,这正是 2.5 要讲的压缩第一路——它直接腰斩这一项。
  • 头维度:每个 K、V 向量有多少个分量。它等于隐藏维度除以注意力头数,是单个头的表示宽度。
  • 精度字节:每个分量占几个字节。FP16 是 2 字节,INT8 是 1 字节,FP8 约 1 字节,更激进的 2-bit、4-bit 量化还能再压。这一项就是"存粗一点"的压缩空间。
  • 序列长 × 批大小:缓存的 token 总数。序列长决定单请求多长,批大小决定并发多少请求。这两项乘积是缓存"随用随涨"的真正引擎。

把这六个因子摆成一列,就能看出哪几项可控、哪几项由模型定死。层数、头维度是模型出厂设定;KV 头数靠架构选择(GQA/MQA);精度字节靠量化;序列长和批大小是运行时的负载。后四项,正是压缩三路和调度策略的发力点。

7B 与 70B 的三档账

给两个有代表性的模型套公式。超参取业界常见配置:7B 按 Llama-2-7B 设(32 层、32 个 KV 头、头维度 128、FP16 精度);70B 按 Llama-2-70B 设(80 层、8 个 KV 头、头维度 128、FP16 精度,注意它用了 GQA,KV 头只有 8 个)。权重以 FP16 计:7B 约 14GB,70B 约 140GB。

模型 精度 上下文 每 token KV 单序列 KV 总占用 相对权重
7B FP16 4K 0.50 MB 2.0 GB 小于权重
7B FP16 32K 0.50 MB 16.0 GB 超过权重
7B FP16 128K 0.50 MB 64.0 GB 约 4.6 倍权重
70B FP16 4K 0.31 MB 1.25 GB 远小于权重
70B FP16 32K 0.31 MB 10.0 GB 小于权重
70B FP16 128K 0.31 MB 40.0 GB 小于权重

这张表的第一条信息:同一精度下,每 token 的 KV 占用是定值,和上下文长度无关,变的是乘法末尾的序列长。所以"长上下文吃显存"不是因为每个 token 变贵,而是因为 token 数量爆炸。

第二条信息,也是最该记住的对比——7B 在 32K 时 KV 已经压过权重,128K 时膨胀到权重的四倍多;70B 因为有 GQA 把 KV 头砍到 8,即便上下文拉到 128K,KV 也只有 40GB,仍不到 140GB 权重的三分之一。同样的上下文长度,架构选型(是否 GQA)能让缓存体量差出好几倍。这给"缓存是回收"补了一句注解:回收的效率一半由你选的模型架构决定。

再补一笔批大小:上面都是单序列。真实服务并发几十上百路,占用要再乘批大小。7B 在 32K、并发 32 路时,KV 就是 16GB × 32 = 512GB——这已经不是单卡能扛的量。所以长上下文 + 高并发是显存的两台抽水机,一起开就见底。

FP16 与 INT8 的取舍

精度字节这一项,是离"存粗点"最近的一刀。FP16 用 2 字节,INT8 用 1 字节,直接把上面整张表的数字对半砍。7B 在 128K 的 64GB,降到 INT8 就是 32GB;70B 的 40GB 降到 20GB。

但这刀有代价。KV Cache 里存的是注意力的键和值,它们直接参与打分和加权平均,精度掉了,注意力分布就会偏移:远距离、低权重的长尾关联最先被抹平,模型在需要"回头精确引用前文"的任务上(如长文档问答、few-shot 里的样例匹配)最容易掉点。FP16 相对安全,INT8 在多数生成任务上几乎无损,再往下到 4-bit、2-bit 就得看具体模型——有的模型对低位量化极敏感,有的靠校准集能救回来。

工程上的经验是:先按 FP16 算账定预算,再用 INT8 当"安全减震",把省下的显存要么换成更长的上下文,要么换成更高的并发。把这事做成动态开关更有价值:短请求用 FP16 保质量,长请求或高并发时切 INT8 换吞吐。下一节要讲的 PagedAttention 恰好为这种"按块灵活处理"提供了基础设施。

一个反直觉的拐点

把整节收束到一个结论:上下文越长,KV Cache 比权重更容易成为显存瓶颈。

直觉上大家觉得"显存主要被模型吃掉",买显卡是为了装下更大模型。但账本摊开看,权重是常量,KV Cache 是随序列长和并发膨胀的变量。7B 这个量级在 32K 以上就已经越过拐点;即便 70B 有 GQA 兜底,长上下文高并发下缓存依然会反客为主。所以"加显存为了跑更大模型"这句话,在长上下文时代要改成"加显存为了缓存更长的对话、扛更高的并发"。

这个拐点不是吓人,而是动机:它解释了为什么需要一个把碎片降到最低的分配器(2.3 的 PagedAttention),为什么需要把相同前缀在请求间共享(2.4 的 RadixAttention),以及为什么要在架构和精度上提前给缓存瘦身(2.5 的压缩三路)。三者都是被这张账本逼出来的。算不清账,就看不见这些设计的价值;算清了,后面几节就是顺理成章。

回到工程现场,这张账本最直接的用处是定"最大上下文"和"最大并发"两个旋钮。给定一张卡的显存,先扣掉权重和激活的刚性开销,剩余能容纳多少 token 的 KV,就框定了单请求上下文上限;再除以计划并发路数,得到每路能摊到的额度。很多服务把上下文上限设得比模型声称的小,根因就是这道除法,而不是模型能力不够。账本还顺手解释了另一个常见现象:同样一张卡,跑 7B 能开很长上下文,跑 70B 却要收紧——不是 70B 算不动,是它的权重把显存基线抬太高,留给缓存的余额被挤薄了。

下面用一个可直接改参数的小计算器,把上面所有数字交到你手里——换个层数、换个 KV 头数、换精度,立刻知道账怎么变。

def kv_bytes(layers, kv_heads, head_dim, seq_len, batch=1, precision=2): # 2 来自 K 和 V 两套矩阵;precision 为每参数字节数,FP16=2,INT8=1 per_token = 2 * layers * kv_heads * head_dim * precision total = per_token * seq_len * batch return per_token, total # Llama-2-7B:32 层,32 个 KV 头(标准多头),头维度 128,FP16 pt7, t7 = kv_bytes(32, 32, 128, 32768, 1, 2) print("7B 每 token KV:", pt7, "字节 =", round(pt7 / 1024 / 1024, 3), "MB") print("7B 单序列 32K 占用:", round(t7 / 1024**3, 2), "GB") # Llama-2-70B:80 层,8 个 KV 头(分组查询),头维度 128,FP16 pt70, t70 = kv_bytes(80, 8, 128, 32768, 1, 2) print("70B 每 token KV:", pt70, "字节 =", round(pt70 / 1024 / 1024, 3), "MB") print("70B 单序列 32K 占用:", round(t70 / 1024**3, 2), "GB") # 切到 INT8,精度字节改为 1,占用直接减半 _, t7_int8 = kv_bytes(32, 32, 128, 32768, 1, 1) print("7B INT8 单序列 32K 占用:", round(t7_int8 / 1024**3, 2), "GB")

运行输出:

7B 每 token KV: 524288 字节 = 0.5 MB 7B 单序列 32K 占用: 16.0 GB 70B 每 token KV: 327680 字节 = 0.312 MB 70B 单序列 32K 占用: 10.0 GB 7B INT8 单序列 32K 占用: 8.0 GB

seq_len 改成 131072、batch 改成 32,计算器会立刻告诉你 7B 在 INT8 下 128K、32 路并发的占用——这正是检验你是否真懂这笔账的小测验:先心算,再跑脚本对账。


作者与出处
原作者: 灏天文库
整理: 灏天文库整理
本站整理收录,版权归原作者/开源协议所有;欢迎通过原文链接访问源仓库。
发布者: 作者: 灏天文库 转发
评论区 (0)
U