KV 缓存与分组查询注意力
自回归生成时每多写一个词就把前文重算一遍太浪费——缓存历史 K/V,再给它瘦身,让推理又快又省。
自回归生成,藏着大量重复劳动
回看逐词生成:每生成一个新词,就把它接到序列末尾,再把 整段 前向一次,得到下一个词。听上去顺理成章,但这里藏着一笔浪费。
能不能把算过的东西 存下来 ,下一步直接取用?
KV 缓存:把历史的 K、V 存起来
在注意力里,新词的 Query 要和 所有 历史词的 Key、Value 做运算。而历史词的 K、V, 不会因为来了新词而改变 。
这一点靠的是学过的 因果掩码 。每个词只看得到它前面的词,新词接到末尾,并不进入任何历史位置的视野,于是历史词在每一层的表示、连同由它算出的 K 和 V,都原样不动。换成双向注意力就没有这条性质,新词会改变历史词在上层的表示,缓存也就无从谈起。
既然历史的 K、V 是固定的,就把它们 缓存 起来。这就是 KV 缓存(KV Cache) :
- 每生成一步,只计算 新词 的 Q、K、V;
- 把新的 K、V 追加 进缓存;
- 让新词的 Q,与缓存里 全部 的 K、V 做一次注意力。
于是每步的注意力计算,从「把已生成的整段重算一遍」,降到「只算新词的 Q、K、V,把新的 K、V 追加进缓存,再让这个 Q 与缓存里全部的 K/V 做一次注意力」。历史那部分的重复劳动被彻底免掉。
还有一点要先交代:注意力层不止一层。每一层都有自己的 Q、K、V 投影,也就各自缓存自己那份 K、V。模型有 32 层,缓存里就有 32 份,互不相干。
下面把这个过程一步步走一遍。上下两栏生成同一句话,上面每步整段重算,下面用缓存,图里画的是其中一层的 K、V。
为什么只缓存 K 和 V,不缓存 Q
每一步真正需要的,只是 当前 这个词的 Query,用它去查询历史。而历史词的 Query,在它们各自被生成的那一步早就用过了,之后 再也不需要 。
被反复使用的,是历史的 Key 和 Value(每个新词都要拿它们来做注意力)。所以缓存的是 K 和 V,名字也因此叫「KV 缓存」。
缓存省了时间,代价是显存
KV 缓存解决了重复计算,但带来一个新负担:
能不能在 不太损失效果 的前提下,减少要缓存的 K、V?突破口在上式的 KV 头数 上。
从多头注意力到分组查询注意力
在「多头注意力」里,每个头都有自己 独立 的 Q、K、V,KV 头数就等于 Q 头数 。而每个头分到的维度是 ,两者一乘正好是 。这说明:在 MHA 里单纯增减头数,缓存量并不会变,头数翻倍,每个头的维度就减半。
要真的省下缓存,得让 KV 头数少于 Q 头数 ,同时保持每头维度不变。 分组查询注意力(Grouped-Query Attention,GQA) 做的正是这件事:保留多个 Query 头 ,但让若干个 Query 头 共享同一组 K/V 头 。Q 的投影维持原样,K、V 的投影矩阵输出维度随 KV 头数一起缩小(每头维度不变,只是头数变少)。
以 LLaMA3-8B 的配置为例,32 个 Q 头、只有 8 个 KV 头, 每 4 个 Q 头共享一组 K/V 。每头维度仍是 128,要缓存的 K/V 从 降到 ,正好是原来的 ,显存大降,而实测精度基本不变。
下面把这套共享结构画出来。KV 头数可以自己拖,从「每个 Q 头各有一组 K/V」一直拖到「32 个 Q 头共用一组」。
MHA / MQA / GQA:一条谱系
把三者按 KV 头数从多到少排在一条线上,区别就清楚了。
共享 K/V 的代价在表达能力上。每个 KV 头给出历史信息的一种编码,K 决定按什么来匹配,V 决定取回什么内容。MHA 里 32 个 Q 头各配一套,能从历史中读出 32 个不同侧面。共享之后,同组的 Q 头仍能算出各自不同的注意力分布,但可供取用的 K、V 只有同一套,输出只能是这同一批 V 的不同加权。KV 头压得越少,各头能提取的信息越受限。MQA 只剩一组,损失最明显;GQA 留 8 组,实测足够把精度保住。
动手:给生成循环加上 KV 缓存
带缓存的生成分成两个阶段,各有名字。第一步把整段 prompt 一次前向、填满缓存,叫 prefill (预填充);之后每步只前向新生成的那一个 token,叫 decode (解码)。两者走同一份代码,区别在一次喂进去几个 token,以及随之而来的掩码与位置偏移。
缓存要在三个层次之间传:注意力层里追加,模型级把每层的缓存收成一份列表,生成循环把这份列表在步与步之间递下去。下面这段代码把这条路径完整走一遍,跑的是一个 4 层的小模型,喂 3 个 token 的 prompt,再生成 3 个词。
缓存省掉的是重复计算,每一步的结果一个都没变,两种写法生成的文本 完全一致 。那到底省下了多少?注意力的计算规则是死的,这笔账可以直接算出来。下面这张图按 LLaMA3-8B 的结构,统计一次前向要做多少次浮点乘加:横轴是序列当前的长度,两条线分别是无缓存与有缓存。
动手:GQA 的形状变换与显存账
把 MHA 改成 GQA,代码上动两处。一是缩小 K、V 的投影,Q 的投影维持 32 头不变,K、V 的输出维度只留 8 个头的份额,切头时也按 8 头切。二是算注意力之前,把每组 K/V 复制 4 份,因为 Q 有 32 头、K/V 只有 8 头,两边对不上。下面这段代码走 decode 的一步,缓存里已经躺着 5 个 token,这一步再往下写一个新词。
既然算注意力前又复制回 32 组,GQA 到底省在哪? 注意力本身的计算量没省 ,复制之后的矩阵乘法规模和 MHA 一模一样。省的是 要存下来、要反复搬运的字节 :复制出的 32 组是当前这一步的中间结果,算完即弃;留在缓存里跨步复用的,自始至终只有 8 组。
缓存量的因子都在手里了,可以真算一笔账。层数 32、每头维度 128 由 LLaMA3-8B 的结构定死,上下文长度、并发条数、缓存精度、KV 头数则可以拨。下面这本账把这些因子乘起来,按 MHA 的 32 个 KV 头与 GQA 的 8 个各算一遍。
账本停在默认那一档时,正文这笔账就出来了:一条 4096 长度的序列,MHA 要 2 GiB,GQA 只要 0.5 GiB。这还只是单条,线上服务同时处理几十条请求,差距要再乘以并发数。省下来的显存可以换成更长的上下文,或者更高的并发。
收益还不止显存容量,生成速度也会提升,不过要先说清瓶颈换了位置。没有缓存时,每步都要把整段历史重算一遍,计算量随已生成长度增长,瓶颈压在算力上,这正是 KV 缓存要解决的。缓存把计算量压下去之后,decode 每步只算一个 token,算力不再是限制,瓶颈随之移到别处。
移到了哪?参数和缓存都躺在显存里,而计算单元不能直接在显存上做乘法,数据必须先搬进来,搬运速度的上限叫 显存带宽 。decode 每生成一步,都要把 模型权重和整份 KV 缓存 搬一遍,只为算这一个 token。时间几乎全花在搬运上,算力大半用不上。这种状态叫 访存受限(memory-bound) ,是 LLM 推理优化绕不开的一个词。
KV 这一块降到原来的 1/4,每步要搬的字节就少一截。但别指望生成也快 4 倍,权重才是搬运量的大头(LLaMA3-8B 的 fp16 权重约 16 GB,而单条 4096 序列的 KV,即便换成 MHA 也只有 2 GiB)。上下文越长、并发越高,KV 在总搬运量里占的比重越大,这份提速才越明显。把上面那本账切到「每步要搬多少字节」,再把上下文和并发拉高,能看到 KV 的占比怎么爬上来。
两笔账合起来看,KV 缓存把每步的计算量从「整段重算」压到「只算一个词」,GQA 再把要存下来、要反复搬运的 K/V 压到 1/4。前者解决算力,后者解决显存与带宽。
零件全齐了
到这里,现代大模型的结构零件已经集齐: RMSNorm、SwiGLU、RoPE、GQA 。 KV 缓存 不带任何参数,不算结构的一部分,它是让训好的模型跑得快的推理手段。
是时候把它们组装成一个完整的现代大模型,并让它真正跑起来了。这是「组装 LLaMA3」要做的事。