☰
大模型推理优化:MHA、MQA、GQA与KV Cache全面解析
2026/10/3 5:50:55 网站建设 项目流程

1. 先厘清背景:从一个面试题和一个部署痛点说起

聊大模型,绕不开这四个词:MHA、MQA、GQA、KV Cache。面试八股里它们是高频考点,实际部署推理服务时它们又是真金白银的显存开销和延迟瓶颈。我在本地部署7B模型、调推理参数、压测并发的时候,对这几个概念的体感尤其深——调不对,要么显存直接爆掉,要么首token慢得让人怀疑人生。

先给没基础的朋友一句话定位:MHA、MQA、GQA是三种不同的多头注意力计算方式,KV Cache是推理阶段用来缓存历史计算结果的机制。它们之间不是并列关系,而是“注意力机制演化”和“工程优化手段”这两条线在推理阶段交汇在了一起。理解了这条线,你就理解了大模型推理为什么快不起来、显存为什么总不够用,以及各家模型架构为什么在这几个方案里反复横跳。

这篇文章不讲虚的,直接把四种机制的来龙去脉、参数计算、选型逻辑和实操配置一次讲透。适合正在学Transformer原理的初学者,也适合已经在部署推理服务、想优化显存和吞吐的工程党。

2. MHA、MQA、GQA:三种多头注意力机制详解

2.1 MHA:标准多头注意力,并行关注多个子空间

MHA全称Multi-Head Attention,也就是标准多头注意力机制。这是Transformer原文里定义的方案,从2017年提出到现在,几乎所有主流大模型的骨干架构都还在用它作为基础组件。它的核心思路不复杂:与其让模型用单一的注意力计算去捕捉序列内所有依赖关系,不如把查询、键、值都投影到多个不同的子空间里,让每个头各自关注不同维度的信息。

具体计算时,输入X先经过三组权重矩阵分别得到Q、K、V。如果头数是h,那么Q、K、V都会被切成h份,每一份对应一个头,每个头独立计算Attention(Q_i, K_i, V_i),最后把所有头的输出拼起来再经过一个输出投影矩阵。公式层面就是标准的:

head_i = Attention(Q_i, K_i, V_i) MultiHead(Q, K, V) = Concat(head_1, ..., head_h) * W_O

为什么MHA效果好?一个直观的理解是:不同的头可以学习到不同类型的依赖关系。比如有的头倾向于捕捉相邻词之间的局部语法关系,有的头能关注到跨长距离的共指关系,有的头可能专门注意句法角色。多个头并行,等于让模型从多个角度同时建模序列关系,表达能力自然更强。

我自己的理解是,MHA本质上是一个“集成学习”的思维——把大模型拆成多个小注意力子空间,每个子空间学到的模式不同,最后拼起来整体效果好于任何一个单独的“大头”。这个思路在视觉领域也被验证过,多头机制配合不同初始化,确实能学到差异化特征。

MHA的代价也很明显:每个头都需要单独维护KV矩阵,头的数量一多,参数量和计算量线性上涨。更重要的是,在推理阶段,每个头都需要缓存自己的KV状态,这为后面KV Cache的显存爆雷埋下了伏笔。

2.2 MQA:所有头共享一组KV,内存直接砍掉大头

MQA全称Multi-Query Attention,思路非常暴力——既然KV占大头,那就让所有查询头共享同一组KV。也就是说,Q仍然分成h个头,但K、V只保留一份。这样在推理阶段需要缓存的KV矩阵直接变成原来的1/h,显存占用断崖式下降。

MQA的做法听起来像是偷工减料,但它是在仔细权衡后的取舍。论文和后来大量实验证明,KV共享对模型精度的影响远小于预期——某些任务上甚至几乎没有掉点。原因可以这样理解:注意力计算中真正需要多样化的是查询Q,因为不同的查询token需要从同一个记忆中提取不同信息;而KV代表的是上下文信息本身,不同头对同一段上下文信息的投影差异,实际上可以通过输出投影层来补偿。

我第一次接触MQA时,觉得这更像是一个“工程上讨巧”的算法。后来看了一些分析文章才意识到,它的成功有内在逻辑——不能让观点不同的“查询者”各自持有完全独立的记忆副本,大家共享同一份信息档案,只是各自提取角度不同。这种设计在规模做大后,能省下巨大的显存和带宽。

MQA的缺点也有,就是当模型层数极深、头数极多时,KV的“瓶颈”效应会更明显,部分任务上表现出稳定性的下降。所以后来才有GQA,在MHA和MQA之间找一个中间地带。

2.3 GQA:分组查询注意力,MHA与MQA的折中

GQA全称Grouped Query Attention,是折中方案。它把查询头分成若干组,每个组共享一份KV,而不是全局共享一份,也不是每个头独立一份。比如32个查询头分成8组,每组4个头共享一份KV,那么KV的数量就是8份,显存占用降为MHA的1/4,精度表现则比MQA更接近MHA。

GQA设计的精妙之处在于,它把“KV头数量”变成了一个超参数,可以在显存占用和模型质量之间做平滑调节。想要更省显存就减少组数,想要更高精度就增加组数。这个特性在实际工程中非常有用,因为不同场景对显存和质量的权衡不一样——长上下文场景对KV Cache的占用极其敏感,这时候就可以考虑减小分组数。

在目前的主流开源模型里,GQA已经成了事实标准。我用的几个代表性模型——LLaMA 2和LLaMA 3、Mistral、Mixtral、Qwen 2系列,基本都采用了GQA。各家选择的组数不一,常见有4组、8组不等。大型模型往往倾向于更大的组数,比如LLaMA 3 70B用了8组,保证质量和速度的平衡。

为什么GQA在精度上比MQA更稳?我个人的理解是:完全共享一份KV,实际上是强制所有查询头从同一个信息空间中提取内容,这会限制不同头学习到差异化的“记忆投影”。而分组后,每个组内共享、组间独立,模型既保留了多样性,又压缩了缓存开销。本质上是把MHA“每头独立记忆”的冗余度压缩到可控范围,同时借用了MQA“共享记忆”的省显存优势,算是“既要又要”的一个实现。

2.4 三种机制的核心差异对比

直接上对比表,方便后续选型时对号入座:

对比维度MHAMQAGQA
KV头数量等于查询头数h1份1组=组数g份
KV Cache显存高(基线)降低约h倍降低约h/g倍
训练/推理计算量高低中
模型质量基线/最好可能轻微下降接近MHA
代表模型Transformer原始论文、早期BERT等PaLM、部分T5变体LLaMA 2/3、Mistral、Qwen 2
适用场景对显存不敏感、重质量显存极端受限场景生产环境主流选择

从这张表能看出来,GQA是兼顾质量和资源消耗的稳妥选择。MQA适合对延迟和显存极端的场景,比如移动端推理、小显存卡部署。MHA则更适合追求极致质量的训练阶段——实际上很多模型训练时依然用MHA,推理时再通过结构转换或蒸馏适配GQA/MQA,不过这属于进阶话题了。

3. KV Cache:大模型推理加速的关键机制

3.1 为什么需要KV Cache:避免重复计算

KV Cache的本质是“用显存换时间”,它解决的核心问题是:大模型推理时,每个token的生成都依赖前面所有token的注意力信息,如果每次都重新计算全部历史token的K、V值,计算量会随着序列长度二次方增长,完全不可接受。

具体到生成流程里,LLM的推理是逐token解码的。生成第n个token时,如果不对KV做缓存,就要从头开始,把整个输入序列重新过一遍模型。比如你输入1000个token,它生成第101个token时要重算这1000个token的K、V,生成第102个时又要重算1001个token。这个过程随着生成长度增加,每秒的计算量会不断膨胀,延迟越来越高,最终导致谁也等不起。

KV Cache的做法是:在第一次计算中把每个层的K、V计算出来并存入显存;后续生成只需对当前最新token计算Q、K、V,然后拿着它的Q去和之前缓存的所有K做注意力计算,再拿新算出的K、V更新缓存。这样每一步的计算量就只和当前token相关,不再随历史长度线性重算。简单说,KV Cache就是“给Transformer配了个备忘录”——第一次查过的资料记下来,后面直接翻笔记就行,不用每次从头翻书。

没有KV Cache,大模型生成一篇几百token的回答会慢到让人抓狂。有了KV Cache,首次计算只需一次前向,后面的每一步都是轻量级的增量计算。这是现代所有LLM推理框架(vLLM、TGI、TensorRT-LLM、Ollama等)的默认底层机制。

3.2 KV Cache到底占了多大显存:一个具体计算实例

KV Cache的显存占用不是小数目,它由层数、KV头数、头维度、序列长度、精度这几个因素共同决定。我直接给一个计算公式,再带入一个真实场景验证。

对于一个Decoder-only模型,单层KV Cache占用公式为:

单层KV字节数 = 2(K和V两份) × 序列长度 × KV头数 × 头维度 × 精度字节数

总占用再乘以层数L。举个例子,假设我用一个7B模型,总共32层,KV头数为8(GQA配置),每个头的维度为128,精度用BF16(每个数占2字节),上下文长度4096,同时处理8个并发请求。计算过程:

  • 单个请求、单个层的KV = 2 × 4096 × 8 × 128 × 2 = 16MB
  • 单个请求、全部32层的KV = 16MB × 32 = 512MB
  • 8个并发请求全部占满上下文 = 512MB × 8 = 4GB

这个数字已经很夸张了——单是KV Cache就把4GB显存吃掉了,还没算模型参数和激活值。如果换用MHA(32个KV头),在同样并发下KV Cache会涨到16GB,一张24GB的消费级显卡基本就告急了。这就是为什么GQA能成为行业标配的重要原因:在大并发、长上下文的真实场景下,KV Cache的显存规模是决定部署可行性的关键瓶颈。

3.3 Prefill与Decode两个阶段

KV Cache的生命周期贯穿LLM推理的两个阶段:预填充阶段(Prefill)和解码阶段(Decode)。

Prefill阶段是模型接收完整输入提示词后,一次性并行计算所有KV Cache并生成第一个token的过程。这个阶段的特点是计算密集度高,因为要同时处理整个输入序列的所有token,GPU利用率很高,但延迟取决于提示词长度。这也是为什么有些框架会把Prefill单独优化,用更大的batch提高吞吐。

Decode阶段则是逐token生成的循环过程——每次只计算当前最新token的Q,用它与已缓存的KV做注意力,生成下一个token,再更新KV。它的特点是访存密集度高,计算量不大但每个token都要反复搬运KV数据,因此GPU算力利用率远低于Prefill。实际上这也是为什么推理提速的核心矛盾不是算力,而是显存带宽——每一步都要把全部KV读一遍。

理解这两个阶段的差异,对调优很有帮助。一个常见误区是:总生成时间=首token时间+每个token时间×(总token数-1)。很多人只看首token延迟,却忽视了Decode阶段逐token生成才是用户体验的大头。实际部署中我通常把这两个阶段分开压测——首token延迟主要受Prefill性能和提示词长度影响,Decode速度则反映模型结构和KV缓存机制的综合效率。

3.4 KV Cache与MQA/GQA的联动关系

KV Cache的大小直接由KV头的数量决定,所以MHA、MQA、GQA的选择会直接作用于KV Cache的显存占用。MHA每个查询头配一份KV,KV头数=h;MQA所有查询头共享一份KV,KV头数=1;GQA每个组共享一份KV,KV头数=组数。

在推理框架里,KV Cache的分配通常按“(最大KV头数 × 层数 × 序列长度 × 精度)”做预分配。使用MQA/GQA时,KV头的实际数量更少,所以预分配显存更小。如果换用MHA,同样的并发和序列长度下,KV Cache占用直接翻几十倍。

这就是为什么很多开源模型设计时会把KV头数显式写在配置里。比如llama系列的配置文件里就有“num_key_value_heads”这个字段——它和“num_attention_heads”不一样。后者是查询头数,前者才是决定KV Cache大小的关键参数。在加载模型或者手动配置推理引擎时,这两个参数经常被搞混,搞错了轻则显存浪费严重,重则直接OOM。我自己的经验是:凡是接触一个新模型的推理配置,首先用官方config.json核对这两个字段,再根据目标并发和上下文长度做KV Cache的预分配规划,不要想当然。

4. 工程实战:如何选择与配置

4.1 不同模型的选型参考

实际选型时,模型架构已经由开源社区定好了,我们能做的更多是理解自己手头模型的KV头配置,然后基于此调优部署参数。我列几个当前常见模型的注意力配置作为参考:

模型注意力机制查询头数KV头数分组数
LLaMA 2 7BGQA32321(实际为MHA)
LLaMA 2 13BGQA40401(MHA)
LLaMA 2 70BGQA6488
Mistral 7BGQA3284
Qwen 2 7BGQA2847
MPT 7BMQA321全部共享

这里有个比较反直觉的点:LLaMA 2的7B和13B版本,虽然官方口径叫GQA,但实际上KV头数等于查询头数,本质上就是MHA。真正启用GQA压缩的是70B版本。所以不能光看模型叫GQA就默认它省显存,一定要去config.json里确认键值头数量。

Mistral 7B的GQA配置很有代表性:32个查询头分成4组,每组8个查询头共享一份KV。这样KV Cache比同体量的MHA模型省了4倍显存。这也是Mistral 7B能在消费级显卡上跑得很舒服的原因之一。

Qwen 2 7B的配置更有意思,查询头28个,KV头4个,分组数7。这説明设计者希望在不牺牲太多质量的前提下极致地压缩KV缓存,也符合Qwen系列在中长上下文场景下的部署需求。

4.2 KV Cache相关的推理参数配置

在实际部署推理服务时,几个和KV Cache紧密相关的参数一定要看懂,不然很容易踩坑:

  • max_batch_size:并发请求数,直接影响KV Cache总量。增大并发必须先预估KV Cache增长。
  • max_seq_len/max_model_len:允许的最大序列长度,其中包含输入和输出token。上下文越长,KV Cache越大。
  • gpu_memory_utilization:框架允许使用的显存比例。vLLM等框架会按这个比例预留显存给模型权重和KV Cache。
  • block_size:vLLM这类PagedAttention框架中KV Cache分配的块大小。块太小管理开销大,块太大浪费显存,一般8或16比较合适。
  • num_key_value_heads:KV头数,决定KV Cache大小的关键参数,一般从config.json读取,不需要手动修改。

我拿vLLM部署一个Mistral 7B模型做示例,配置如下:

model = "mistralai/Mistral-7B-Instruct-v0.3" tensor_parallel_size = 1 max_model_len = 8192 gpu_memory_utilization = 0.85 enforce_eager = False

这里把最大长度设为8192,显存利用率85%。在24GB显卡上,模型权重约14GB,剩下约6.5GB留给KV Cache和激活值。实测并发8个请求、每个请求最长生成1024 token时,系统运行稳定。如果我把最大长度调到16384,KV Cache预分配会翻倍,显存立刻吃紧,并发只能降下来。这就是KV Cache和并发、上下文长度三者之间的三角博弈。

4.3 显存不足时的优化策略

显存不够怎么办?我按自己的优先级列表排了一个实战顺序:

第一优先级:确认模型是否启用了GQA/MQA。如果模型本身是全MHA架构且对质量不是极端敏感,可以尝试切换到MQA/GQA变体或微调版本,显存压力能直接降一个量级。

第二优先级:降低最大序列长度。很多时候业务场景并不需要极长上下文,把max_model_len从8192降到4096,KV Cache直接少一半。这个操作最立竿见影。

第三优先级:降低并发或者开启请求排队机制。通过控制并发上限,避免多个请求同时占满KV Cache。vLLM的continuous batching会动态调度,但峰值并发仍然受资源上限约束。

第四优先级:使用KV Cache量化。把KV Cache精度从BF16降到INT8甚至INT4,显存再降一半甚至更多。代价是精度损失可能影响长文本质量,在推理框架里属于可开关的高级特性。

还有一个容易被忽略的点:是否显式开启use_cache=True。某些库默认会关闭KV Cache,导致生成速度慢到不可接受,还以为模型有问题。第一次踩这个坑时,我排查了很久才发现是默认参数搞的鬼。

5. 常见疑问与踩坑记录

5.1 为什么GQA的分组数一般是8而不是其他数?

这个问题的答案和显存、性能、质量三者的平衡有关。分组数越大,KV头越多,显存占用越高,但精度越接近MHA;分组数越小,KV头越少,显存越省,但可能出现KV信息不足、模型质量下降。8这个数值在大多数任务上已经能逼近MHA效果,同时显存压缩比例达到h/8,收益和代价都处于甜蜜点。比如70B模型32个查询头,分8组每组4个共享一份KV,KV数量缩小为8份,压缩4倍,而质量损失基本可以忽略。

当然也有特例。Qwen 2 7B用了7组,Mistral 7B用了4组,这体现了不同团队对模型质量和效率权重的主观取舍。如果应用场景非常在意长上下文记忆的细节还原度,可以考虑少压缩一些;如果只是短对话场景,可以选MQA或者少组GQA进一步压显存。

5.2 MQA真的会掉点吗?

MQA掉不掉点,不能一概而论,要在足够大的模型规模上才能下结论。论文数据和很多实践表明,在小模型上MQA相对MHA的掉点比较明显,但只要模型规模上去,MQA和MHA的差距会缩小到接近噪声水平。

我自己的猜测是:小模型本身容量有限,信息存储空间不足,让所有查询头共享同一份KV,容易造成信息瓶颈;而大模型的参数多、表示能力强,模型可以通过其他层来补偿老KV信息多样性的缺失。所以如果你的场景是小模型落地,选择MQA要谨慎,先跑一批评测数据对比;如果是大模型(10B以上),MQA作为极速推理方案是值得尝试的。

另外,训练时用MQA和训练后强制改MQA是两回事。MQA最好在预训练阶段就确定;用MHA预训练好的模型,推理时强行改成MQA结构,通常需要额外的蒸馏微调才能恢复大部分质量。这一点容易被忽略,很多人说“MQA掉点严重”,往往是在后者的情况下得出的结论。

5.3 实测压测中的几个细节

本地部署推理服务时,有几个细节很容易导致压测结果失真:

第一,框架的KV Cache复用机制要开启。vLLM支持prefix caching,对共享前缀的请求能复用KV缓存,显著提升吞吐。我压测时发现,同样是并发生成,请求前缀完全相同时的吞吐是随机前缀的几倍。这直接影响容量规划,不要拿理想案例当容量上限。

第二,显存碎片问题。小请求不断申请和释放KV Cache,可能造成显存碎片。PagedAttention把KV Cache按block管理,能缓解这个问题,但block_size设置不合理也会白白浪费显存。

第三,测延迟和测吞吐要分开。延迟优化要压低并发、拉长batch超时时间;吞吐优化则要抬高并发、动态调度。两者对KV Cache的需求曲线不同,混在一起测会得到完全看不懂的结果。

第四,量化KV Cache是否开启要实测。INT8量化KV Cache在某些模型上有轻微质量下降,但在显存紧张的场景下,降量化开启后能跑的并发量可能提升最高一倍。建议用评测集跑一遍关键指标,自己判断深浅程度是否能接受。

最后再分享一个小细节

如果你在读模型源码或者调试推理框架,碰到KV Cache相关的问题,建议先看模型配置里的num_key_value_heads字段。很多看似奇怪的行为——比如某个模型生成的显存占用比预期高好几倍、或者并发稍高就OOM——根源都是KV头数配错或者KV Cache预分配策略没对齐。

我在本地跑一个早期MHA架构模型时,曾因为没搞清楚KV Cache原理,以为显存OOM是模型权重占满,拼命换小模型版本,后来才发现是并发+长上下文让KV Cache爆了。调整并发上限后,同样的模型在同样显存下稳定跑了起来。这几个概念不是只存在于面试题里,它们是AI工程落地时最实在的功课。希望这篇梳理能让你少走一次弯路。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询