资讯动态

自回归生成、KV Cache与GQA:大模型推理优化主线全解析

发布时间:2026/9/8 14:39:54 来源:尧图企业网站定制
1. 开场把Decoder这条主线在大脑里画出来Decoder在Transformer里算是个又熟悉又陌生的角色。熟悉是因为只要聊大模型、聊文本生成就一定绕不开它陌生是因为它身上同时挂着好几套概念——自回归、Teacher Forcing、KV Cache、MQA/GQA、Causal Mask单独拆开每个都有人讲但很少有人把它们串成一条线讲清楚。这第9课我想干的事比较朴素把这些点全部串起来让你从头到尾走一遍Decoder在推理时到底是怎么工作的、每一步优化是在解决什么问题。这课适合两类人。一类是正在啃大模型源码、想搞清楚LLaMA或者GPT模型里那些缓存变量到底在干嘛的初学者另一类是已经会调接口、但是想往模型底层再深入一步想知道GQA这种结构为什么能省显存、省了之后质量到底有没有损失的工程师。课程目标只有一个读完以后你再看到Decoder相关的任何文章或源码脑子里能自动把概念映射到一条清晰的主线上——自回归生成是需求KV Cache是性能优化GQA是KV Cache的瘦身手术。这三件事本质上是一条因果链。先提前说一句这节课会涉及一点数学和代码但都是最朴素的那种不需要你有多深的功底。我会尽量用“账本”“食堂打饭”这类生活场景去类比保证你听到的不是一堆术语轰炸而是真正能落到代码里的理解。2. 主线的源头自回归生成到底在干什么2.1 自回归的本质一步一步猜下一个词自回归Autoregressive这个词听起来吓人其实拆开就是“自己回归自己”。放到文本生成场景里它的含义很直白模型在生成第 t 个词的时候只能看到已经生成的前 t-1 个词然后根据这些词去预测下一个词的概率分布。生成完一个词把它拼到已有序列后面再继续预测下一个。整个过程像一个人在玩接龙你看到“今天天气”就会接“不错”看到“今天天气不错”可能会接“我们出去走走”。每一步都只依赖已经写下的内容。写成概率公式就是P(y_1, y_2, ..., y_T) P(y_1) * P(y_2 | y_1) * P(y_3 | y_1, y_2) * ... * P(y_T | y_1, ..., y_{T-1})整体句子的概率被分解成每一步的条件概率乘积。这个分解就是自回归的数学根基。Decoder干的事就是一步步把这堆条件概率算出来然后挑一个词作为下一步的输入。为什么要用这种方式因为语言本身就是时间序列——词与词之间有先后依赖先说什么后说什么直接决定语义。自回归天然贴合语言的这个属性。虽然现在也有非自回归模型试图并行生成所有token但质量、灵活度上始终和自回归有差距所以主流生成模型几乎清一色是自回归架构。2.2 训练和推理的“人格分裂”——同一个Decoder两副面孔这是很多初学朋友最容易卡住的地方训练和推理虽然用的是同一个模型权重但输入方式完全不同。训练阶段用的是Teacher Forcing。意思是说不管模型上一不步预测得对不对我都把真实的、正确的词喂给它当下一步输入。比如训练样本是“我喜欢吃苹果”模型第一步输入“我”我强制要求它输出“喜”的概率最大到了第二步不管它输出什么我都强行把“喜”作为输入喂进去让它预测“欢”。这样做的目的是加速收敛——让模型每一步都在学习“在正确上下文下预测下一个词”而不是像推理那样一步错、步步错。推理阶段就没有Teacher了。模型只能把自己上一步的输出当作下一步的输入。这就是“自回归”和“Teacher Forcing”之间最本质的差别。这里有一个关键矛盾训练时模型看到的都是真实上下文推理时看到的却是自己生成的内容这中间有分布偏移叫Exposure Bias会让模型在推理时积累误差。很多缓解手段比如Scheduled Sampling就是为了减小这个偏移。这个差异直接引出了性能问题推理时每生成一个新token计算量是不是在疯狂增长这就为KV Cache埋下了伏笔。2.3 自回归推理的计算浪费在哪里假设我们生成一句话长度是20个词。如果用最笨的办法每一步都把当前完整序列从头算一遍注意力会发生什么第1步输入只有1个token算一次注意力轻松。第5步输入有5个token重新对5个token做一次注意力计算。第10步输入10个token再做一次10长度的注意力。问题来了第5步计算第1个token的Query、Key、Value时和第4步计算第1个token的Key、Value结果完全一样。因为输入的第1个token从头到尾就没变过模型的权重也没变所以它的K和V必然一样。但你在第5步又老老实实重新算了一遍。这就是纯粹的重复劳动。更具体一点假设序列长度是 n每个token的K、V向量维度是 d。计算一次完整序列注意力的复杂度是 O(n² · d)其中很大一部分开销就花在了重复计算历史token的K和V上。每一步都重新对所有历史token做矩阵乘法推理越来越慢计算量呈平方级增长——这就是没有KV Cache的暴力方案。有没有办法把这些重复计算的结果存下来下次直接取用有这就是KV Cache。3. KV Cache把算过的账记在本子上3.1 KV Cache到底缓存了什么KV Cache的思路特别简单既然历史token的Key矩阵和Value矩阵是确定的、不会变的那我第一次算完就把它们存下来之后每次生成长度1只需要计算新token的K和V然后把它们追加到缓存里别的层直接去缓存里查就行。用一个生活类比你每天从家步行去公司会经过一个路口。你第一次走这条路时需要看地图确认方向但走了一百遍之后你根本不需要重新看地图身体已经记住了。KV Cache就是这个“身体记忆”——已经走过的路径不用再重新探索重复利用就好。具体到代码层面KV Cache通常是一个列表长度为模型的层数。每一层里存着两个大Tensor一个是所有历史token的Key矩阵形状是[batch_size, num_heads, cache_len, head_dim]另一个是Value矩阵形状一样。每次生成一个新token模型计算出这个token在当前层的K和V然后拼接到缓存张量的第3维序列长度维度上形成新的缓存。这里有个容易混淆的点为什么只缓存K和V不缓存Q因为Q是当前要预测位置的查询向量每次只算当前这个token的Q没有“历史”可存。而K和V是历史token的身份信息一旦确定就不会再变。3.2 有缓存和没缓存的计算量对比我们来算一笔账。假设模型有 L 层每层有 H 个头每个头的维度是 d当前已经生成的序列长度是 n。没有KV Cache时生成第 n1 个token模型要对全部 n 个token重新计算一次注意力。单层单头的一次注意力分数矩阵是[n, n]需要做 n×n 次乘法。生成一个token的总计算量大约是L × H × n² × d这里忽略FFN等其它部分注意力是大头。有KV Cache时生成第 n1 个token只需要计算新token的 Query 和 Key、Value然后用这个 Q 去和缓存的全部 K 做点积。注意力分数矩阵从[n, n]变成[1, n]——只算一行。计算量大约是L × H × n × d。两者一对比复杂度从 O(n²) 降到了 O(n)。生成的第 n 个 token 越靠后省下的计算量越大。这也是为什么实际推理时KV Cache几乎是标配没有它长文本生成根本跑不动。3.3 KV Cache的显存估算7B模型到底要吃多少KV Cache省了计算但它不是免费的——它要占显存。而且占得非常夸张。直接给一个公式大家以后可以自己估算KV Cache大小 2K和V两份 × L层数 × H头数 × d每头维度 × n序列长度 × batch_size × 每个参数字节数举个例子一个7B模型假设是32层32个注意力头每头维度128生成4096个tokenbatch size是1。硬算一下2 × 32 × 32 × 128 × 4096 × 1 1,073,741,824 个参数如果模型用FP16存储也就是2字节KV Cache就需要约2GB显存。注意这还只是单条序列、单个batch。如果batch size开到32那就是64GB显存直接把一块A100的显存吃干抹净。所以KV Cache这把刀在推理延迟上是利刃在显存占用上是猛兽。这就不难理解为什么后来会有MQA、GQA——它们都是奔着给KV Cache减重去的。3.4 PyTorch里的实现要点在PyTorch或者HuggingFace的源码里KV Cache的实现方式比较统一。以标准Decoder层为例伪代码大概是这样的# 假设 cache 初始为 None # 第一次调用时创建空缓存 if past_key_value is None: past_key_value (torch.empty(0), torch.empty(0)) # 计算当前token的 K 和 V k self.k_proj(hidden_states) # [batch, 1, heads, head_dim] v self.v_proj(hidden_states) # 拼接历史缓存 k torch.cat([past_key_value[0], k], dim2) # 在序列长度维度拼接 v torch.cat([past_key_value[1], v], dim2) # 计算当前token的Q并与拼接后的K、V做注意力 q self.q_proj(hidden_states) attn_weights torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(head_dim) attn_weights attn_weights.masked_fill(causal_mask, float(-inf)) attn_weights torch.softmax(attn_weights, dim-1) output torch.matmul(attn_weights, v) # 更新缓存 past_key_value (k, v)有几个实现细节需要注意。第一个是拼接操作PyTorch的torch.cat在每步生成时会复制整个缓存矩阵序列越长开销越大这也是为什么有像PagedAttention这种针对KV Cache的内存管理方案。第二个是Causal Mask在使用缓存时因为我们已经把历史token都拼进去了新token的注意力理论上能看到所有历史token但绝对不能看到未来的token本来也没有。不过这里藏着一个坑我把它放到后面的常见问题部分细讲。4. GQA给KV Cache做一次瘦身手术4.1 MQA先试一条极端的路KV Cache的显存占用和注意力头数直接相关。自回归模型刚起步时用的是MHAMulti-Head Attention也就是每个Q、K、V都有自己独立的一组头。以7B模型为例32个头每个头的K、V都要存一份显存自然大。有人就想能不能让KV的头少一点让所有Q头共享同一组K和V这就诞生了MQAMulti-Query Attention。MQA的做法很激进不管你有多少个Q头K和V都只有一个头。所有Q都去和这同一个K、同一个V做注意力然后输出再接一个投影层恢复维度。MQA的显存收益非常直观。还是刚才那个7B模型的例子假设32个Q头如果K、V都只有1个头KV Cache的显存直接降为原来的1/32。原来需要2GB现在只需要64MB。这个收益是巨大的。但问题也出现在这里K和V只有单个头意味着所有Q头共享同一份“信息表达”。不同Q头理论上应该关注不同维度的语义比如有的头负责语法有的头负责指代消解共享K/V会限制每个头的表达空间。实测下来MQA在部分任务上会有轻微质量损失尤其在需要细粒度语义理解的任务上。4.2 GQA折中的最优解GQAGrouped Query Attention是Google在2023年提出的方案它站的位置正好在MHA和MQA中间。用大白话说既然每个Q头单独配一组KV太贵全部Q头共享一组KV又太粗糙那我就把Q头分成几组每组共享一组KV。假设模型有32个Q头GQA可以设置成8组每组4个Q头共享一组KV。这样KV头的数量是8个比MHA少4倍KV Cache也缩小为原来的1/4。如果极端一点设成1组GQA就退化成MQA如果设成32组每组1个头GQA就退化成MHA。可以说GQA是MHA和MQA之间的一个可调节旋钮想省多少显存由你决定。这里有个特别直观的类比食堂打饭。MHA是每个厨师做一份菜32个厨师同时开工菜多但贵MQA是一个厨师做一份大锅菜所有窗口都打这份便宜但众口难调。GQA则是把厨师分成8组每组负责一个菜系兼顾了口味和成本。4.3 GQA的PyTorch实现不少朋友在Transformer代码里想改GQA不知道从哪下手。我给一个最简洁的核心实现思路重点是K、V如何从若干个头“扩展”回Q的头数。def repeat_kv(kv, num_groups, num_heads): # kv 形状: [batch, num_kv_heads, seq_len, head_dim] # 目标是每个 Q 头都有对应的 K/V batch, num_kv_heads, seq_len, head_dim kv.shape kv kv[:, :, None, :, :].expand(batch, num_kv_heads, num_groups, seq_len, head_dim) return kv.reshape(batch, num_heads, seq_len, head_dim)在注意力层里Q的头数是num_headsK、V的头数是num_kv_headsnum_groups num_heads // num_kv_heads。计算Q和K的点积前把K以及V沿着头维度重复num_groups次让K的头数对齐Q的头数后面的计算就和标准MHA一模一样了。这样改起来不需要动太多代码只需要在K、V投影的输出通道数上做文章。具体到LLaMA的实现源码KV Cache里存的是压缩后的、头数较少的K和V而不是重复后的版本。这样缓存占用的显存还是按num_kv_heads算只是在每一步注意力计算前临时把K、V重复展开到和Q一样的头数。这个顺序很关键别弄反了。如果你在缓存里存的就已经是重复后的完整张量那GQA省显存的效果就完全没达到等于白干了。4.4 模型实测MHA、MQA、GQA到底怎么选从实际落地的角度来看模型里到底用哪种注意力业界早就有了一批参考数据。LLaMA 2的7B和13B版本用的是MHA但70B版本用了GQA分组数是8。LLaMA 3全系列用的都是GQA而且8B和70B都做了不同组数配置。Mistral 7B用的则是MQA。这其实反映出一个行业共识大模型越往后做越倾向于用GQA因为它提供了一种“用少量质量损失换取巨大显存收益”的可控权衡。为什么大模型更依赖GQA因为模型参数量越大层数越多、头越多KV Cache的膨胀速度越快。7B模型只是起步70B模型的KV Cache如果不做压缩单条长文本就能吃掉几十GB显存这种成本在工程上根本无法接受。GQA的价值在于它把KV Cache的显存占用降低了几倍到几十倍对模型质量的影响却很小——实际测评中GQA在大多数任务上几乎无损只有少数细粒度理解任务会有可感知的回落。所以如果自己做技术选型我的建议是小模型、序列短优先MHA简单直接效果好序列长、并发高、显存紧张果断上GQA分组数可以先取8再根据效果微调只有极端追求速度且能忍受质量轻微下滑的场合才去考虑MQA。5. 主线串联与实操心法5.1 把整条链串起来走到这一步我们可以把前面所有概念串成一条完整的因果链了。自回归生成要求模型逐token预测每生成一个新token都要以全部历史上下文为条件。于是推理时如果不加额外优化就要不断地对历史token重新计算K和V造成大量重复计算。KV Cache解决了这个重复计算问题代价是显存占用随序列长度线性增长。显存一紧张KV Cache又成了瓶颈于是MQA出现了它把KV头数压缩到1个但显然太激进GQA就出来做了折中按组共享KV在显存收益和模型质量之间取得平衡。这就是标题里说的“自回归、KV Cache与GQA引出的完整主线”。这三句话不是三个孤立知识点而是一条解决问题的链路需求决定了架构架构暴露了瓶颈瓶颈催生了优化方案。如果你能把这个逻辑链条刻在脑子里以后再看到任何新的注意力变体——比如MLA、DeepSeek的Multi-head Latent Attention——你都能猜到它大概是在优化哪一环。MLA就是进一步把KV Cache压缩到更低维度的方案本质还是在解决同一个问题。5.2 实操心法一Causal Mask和KV Cache的配合坑这个坑我调试的时候踩过值得单独拿出来说。在没有KV Cache的完整序列训练场景里Causal Mask是一个上三角矩阵用来屏蔽未来位置。但在推理阶段如果你已经用了KV Cache序列是逐步增长的这时候你算注意力实际上只需要当前token的位置去attend所有已有的历史位置不存在“未来token”所以理论上不需要Causal Mask。但很多框架为了代码统一在推理时依然会保留mask逻辑。问题来了如果mask的shape写死了比如原本是[batch, 1, seq_len, seq_len]而随着序列增长seq_len在变如果mask没有跟着扩展到当前长度就会报错或者出现掩码错位。更隐蔽的问题是有些人在拼接缓存时把mask也拼接了但mask拼的是0还是负无穷一旦搞错模型生成质量会莫名下降。我的建议是推理阶段只要保证没有任何未来token进入注意力范围mask可以直接不用省掉一次张量计算。如果你的框架非要传mask请确保它和cache的序列长度严格对齐别让mask的维度滞后或者超前。5.3 实操心法二Beam Search和KV Cache怎么配合Beam Search是生成质量更高的解码策略但它和KV Cache配合时有个容易忽略的问题beam search每一步都会维持多个候选序列比如beam size4每个候选序列的KV Cache是独立的不能混着用。有的实现为了省代码在beam search过程中没有正确维护每个beam自己的缓存结果一个beam的KV被另一个beam复用生成出来的句子语义乱七八糟。正确做法是每做一次beam扩展就要根据新的beam索引对KV Cache做一次gather操作把对应位置的缓存取出来。HuggingFace的源码里这块逻辑写得比较隐蔽但如果你自己实现beam search务必记得这一步。还有个小技巧beam search时如果某个beam被提前终止它对应的KV Cache释放后显存并不会立刻归还给CUDA。长时间运行大batch beam search时需要留意显存碎片化的问题必要时手动清理或者用一个缓存池管理。5.4 实操心法三一个常见的性能误区很多人以为加了KV Cache延迟就一定降低。对但前提是序列已经足够长而且注意力计算占主导。如果是短文本生成比如序列长度只有十几二十个tokenKV Cache拼接本身的拷贝开销可能比重复计算的成本还高——尤其是Python层面调用torch.cat时每步都在显存里搬数据。实测下来序列长度在几十个token以内时KV Cache带来的收益不明显到几百、上千token时优势才真正放大。所以如果你在做一个极短文本生成的低延迟服务可以考虑在开始时不用KV Cache或者直接用torch.cat之前先给缓存张量预留足够的连续空间避免反复申请内存。这个问题在后端工程化时很常见也容易被只看论文的人忽略。5.5 实操心法四GQA分组数的选择不是拍脑袋GQA里的分组数设多少既不是越大越好也不是越小越好它直接决定KV Cache的压缩率和模型的表达容量相互制约。拿7B模型举例假设32个Q头。如果分组数取8KV Cache减为原来的1/4模型质量几乎不降如果分组数取2KV Cache减为原来的1/16质量会有轻微下降如果分组数取1就是MQA这时候质量下降就比较明显了。我的经验是第一次上手可以先从8组开始实验对比一下困惑度perplexity和下游任务指标如果损失不明显再往小调。千万别一上来就图省显存设成1组除非你的任务本身对生成质量要求不高比如粗粒度的关键词提取。另外提醒一句GQA的KV头数最好能整除Q头数否则repeat_kv的时候会很难受。保持8组、4组这种整齐的数字代码写起来也干净很多。6. 常见问题与排查速查这里把这几年带团队、看源码、做推理优化时大家最常踩的问题整理成一张速查表按问题现象、原因、解法来组织方便你对照排查。问题现象根本原因解决办法生成结果随机性大且质量差Causal Mask维度没跟上KV Cache长度检查mask是否与cache的seq_len严格对齐或推理阶段直接省去mask显存随生成token数爆炸式增长KV Cache的保存方式错误可能存了展开后的重复KV确认缓存存的是压缩前的KV按 num_kv_heads 存储注意力计算时才临时repeat多个beam的生成内容互相干扰Beam Search没有按beam索引gather各自KV Cache每次beam扩展后根据新索引对KV Cache做gather使用GQA后模型输出质量明显滑坡分组数设太小或者Q头数无法整除KV头数调大分组数保证num_heads % num_kv_heads 0短文本生成延迟不降反升KV Cache的cat拷贝开销大于减少的计算量短文本场景评估是否值得用KV Cache或预留缓存空间避免反复malloc多轮对话历史越长越慢上下文不断累积KV Cache线性变大考虑滑动窗口注意力、token压缩或状态缓存方案这条速查表覆盖了从训练到推理、从MHA到GQA的大部分实战问题。你可以把它当成一个排查的起点遇到对应现象的时候先按表里的思路去定位大部分问题都能在代码里找到根因而不是在模型质量上反复折腾调参。最后说一点个人体会。KV Cache和GQA这套东西第一次接触时可能会觉得是纯工程技巧不涉及什么高深理论。但后来你会发现大模型从学术产物走向工业落地拼的恰恰是这种“看似朴素、实则关键”的工程优化。每一层缓存、每一个分组数背后都是成本、速度和质量之间的妥协。理解了这条主线不只是看懂几行源码更是理解了大模型推理系统设计的核心思路。这套思路你以后去看任何新出的推理优化方案都能一眼看出它动的是哪一环的奶酪。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价