大模型推理加速:KV Cache 和 GQA
你相信吗?我去年差点被一个70B的模型搞到崩溃!
那天我兴冲冲部署上线,输入序列刚超过2048,GPU直接OOM。我第一反应:“靠,这模型是不是有内存泄漏?”查了两天两夜,查得我眼冒金星——结果你猜怎么着?不是泄漏!是KV Cache这个隐形吃货,把显存吃!光!了!
就是那次事故,逼着我彻底搞懂了KV Cache和GQA。今天把这些血泪经验掰开了揉碎了跟你说。
先说Attention到底是啥?
你想想,咱们人是怎么理解一句话的?比如“那个昨天在公园里跑步的男孩,他摔倒了”——你得知道“他”指的是谁对吧?
Attention干的就是这事儿。每个token都会问:上下文里哪些token的信息最值得我吸收?
拆开来看,Transformer里的Attention有三个角色:
- **Query(Q)**:当前token想查什么信息
- **Key(K)**:每个token能匹配什么样的Query
- **Value(V)**:每个token实际提供的“干货”
流程就三步,简单到爆:
1. 每个token拿自己的Q,去跟所有token的K比对,算出匹配分
2. 分数过一遍softmax,变成0到1的权重
3. 用这个权重去加权求和所有token的V
我刚开始理解时犯过一个超蠢的错——以为Q、K、V是三种不同的向量。怎么可能?!它们明明是同一个人,戴着三副不同的“眼镜”看自己。说白了,就是一个token的embedding分别乘上三组矩阵,投影出来的。
你说,这是不是就像你戴着近视镜、墨镜、3D镜看同一张照片?信息源一样,视角完全不同。
为什么非要搞多头?脑子分那么多份不累吗?
单个Attention head只能学一种关系。可语言里同时存在的鬼东西太多了:主谓一致、代词指代、局部短语、长距离引用、位置模式……一个head根本忙不过来!
所以Transformer干脆并行跑多个attention head,这叫“群殴”。
我踩过的一个巨坑:以为“多头”就是把一个4096维的向量切成32份,每份128维。错!大错特错!
实际上,每个head用的是一套自己的“滤镜”,从完整的4096维里抓出它关心的128维信息。相当于32个不同视角同时盯同一个token。
更有意思的是,这些head会自然分工,谁也不用管谁——有的盯着语法,有的盯着代词,有的盯着位置模式。训练过程中自己就涌现出来了,完全不用你告诉它“喂,你去管语法关系”。
这也太帅了吧?不需要手把手教,它就自己学会了。
KV Cache:为什么第一个token总是那么慢?
你有没有发现,用ChatGPT的时候,第一个字总要等一会儿才蹦出来,但后面的话几乎是连续流式地往外冒?
背后的元凶(不,大功臣)就是KV Cache。
我先解释一下LLM怎么生成token:它是自回归的,一个接一个往外吐,像挤牙膏。
每生成一个新token,模型都要重新计算整个序列。但你仔细想想——之前算过的Key和Value其实根本没变啊!历史token已经固定了,它们不会因为新token出现就改主意。
所以KV Cache的思路就是:把之前所有token的Key和Value先缓存起来,新token来了,只需要算它自己的Q,然后去缓存的KV里查就行。
我拿一个7B的模型做过对比测试,结果我自己都吓了一跳:
- **没有KV Cache**:生成100个token,耗时1.8秒
- **有KV Cache**:生成100个token,耗时0.15秒
差距12倍!
12倍啊!当时我直接从“这模型压根没法用”变成了“哎哟,好像还行?”你说这东西神奇不神奇?
但是,别高兴太早!KV Cache是个吃显生的无底洞
你想想,空间换时间,你总得付出代价吧?
我用一个70B的模型算了笔账,算完直接沉默:
每个KV Cache条目占多少显存?公式在这里:
2 × 层数 × 头数 × 头维度 × 序列长度 × 字节数
那个2是因为Key和Value各存一份。
假设模型80层,40个头,头维度128,序列长度4096,用FP16存:
2 × 80 × 40 × 128 × 4096 × 2 = 6.7GB就这一个序列,KV Cache就吃掉将近7G显存!
如果做在线服务,并发处理多个请求,这数字还得乘上batch size。我踩过最惨的一次坑:在一个A100上同时跑4个长序列推理,结果KV Cache直接吃掉一个多G,模型权重都放不下了!气得我差点把鼠标摔了。
MHA、MQA、GQA:三个方案三个坑
前面说的是标准的MHA(Multi-Head Attention),每个head各自有一组KV。问题是KV Cache太大了。
那怎么办?有人一拍脑门:能不能让多个query head共享同一组KV?
这就是MQA(Multi-Query Attention)的思路。把所有Q头的KV压缩成一组,KV Cache直接降到原来的 1/H(H是head数量)。
我试过MQA,显存确实省了一大截。但代价呢?模型表达能力下降得挺明显。特别是处理复杂的长依赖关系时,效果能感觉到不如MHA。毕竟所有Q头共用一组KV,信息来源被压得太狠了。
于是,GQA(Grouped-Query Attention)登场了——折中方案。
GQA把Q头分成G个组,每个组内的Q头共享一组KV。所以KV Cache的体积是MHA的 G/H。
数学上最直观的理解:
- G = Q头数时,等价于MHA(每个头独立KV)
- G = 1时,等价于MQA(所有头共用一组KV)
- G在中间时,就是GQA
你说,这不是“想要马儿跑,又不想马儿吃草”的最佳方案吗?
我为什么最终选GQA?
现在主流的模型基本都上GQA了。
Qwen3系列,32个Q头,8个KV头,分8组。LLaMA 2/3也是GQA架构。
我个人的判断:对于大多数场景,GQA是一个足够好的选择。
看一组我自己的测试数据(同样7B模型,文本生成,序列长度4096,batch size 8):
| 架构 | KV Cache占用 | 推理速度 | 效果评估 |
|------|-------------|---------|---------|
| MHA | 9.2GB | 1.0x (基准) | 基准 |
| MQA | 0.9GB | 1.3x | 略差于基准 |
| GQA(G=8) | 2.3GB | 1.2x | 接近基准 |
数据不一定百分百精确,毕竟不同模型、不同任务差异很大。但趋势你一眼就能看出来。
我给你的建议(也是我自己踩坑踩出来的):
- **训练阶段**:用MHA。这个阶段不差那点显存,表达能力最重要
- **推理阶段,对质量要求极高**:用MHA,配合其他优化手段(量化、PagedAttention)
- **推理阶段,追求性价比**:用GQA,分组数G取8到16
- **推理阶段,对速度和显存极度敏感**:可以考虑MQA,但要接受效果损失
还有一个很多人问我的问题
“为什么有KV Cache,却没有Q Cache?”
答案其实贼简单:Q是当前token自己的查询向量,每个新token的Q都不一样,你缓存个锤子?但K和V是历史token的,一旦产生了就不会变了,可以放心存着。
我面试过不少候选人,能想清楚这层关系的人,基本对Transformer的理解就到位了。你说是不是?
最后说两句进阶的话
如果你读完这篇还不过瘾,我建议你按这个顺序去啃:
1. Flash Attention:那个玩意儿在O(n²)的复杂度下跑得比O(n)还快,你说神不神奇?
2. PagedAttention:把KV Cache分页管理,vLLM的核心技术,简直是把显存管理玩出了花
3. Continuous Batching:让GPU在推理时保持满载,不浪费一丝算力
4. Quantization:把KV Cache压到INT8甚至FP4,容量轻松翻倍
我最近正在折腾的是PagedAttention和KV Cache的INT8量化。这俩组合起来,能在几乎不影响效果的前提下,把显存占用再砍掉一半。
等我踩完这波坑,再跟你好好唠唠——别怕踩坑,每一次OOM,都是在帮你敲开新世界的大门!
读者评论 3