ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

一文讲透MHA、MQA、GQA与KV Cache:大模型推理加速的核心

一文讲透MHA、MQA、GQA与KV Cache:大模型推理加速的核心 用任何大模型做自回归生成都会碰到一个很有意思的现象你输入一段很长的prompt模型读题阶段基本是并行计算的通常几百毫秒就能出第一个token但紧接着的每一个输出token反而要几十毫秒甚至更久“读题快、出字慢”特别明显。很多第一次接触LLM推理的人会困惑为什么越往后生成越慢答案就藏在KV Cache里。但要真正讲清楚KV Cache绕不开MHA、MQA、GQA这套多头注意力机制的演化史。MHA是Transformer自带的“标准多头”MQA是KV共享的“激进瘦身版”GQA是当前开源大模型最主流的“折中方案”KV Cache则是连接它们与推理性能的核心枢纽。它们之间的区别表面看是参数量和算力差异本质上是显存、带宽、表达能力三者的三角平衡。这篇文章把四个东西放在一起拆开讲从数学形式到工程实现从显存计算到部署选型适合正在学大模型原理、准备面试八股、或者做推理部署时被OOM折磨过的人。1. 固定Q多头、共享KV单头的“瘦身”实验——MQA的诞生背景1.1 从Self-Attention公式看“多头”到底多了什么先回到最基础的Self-Attention公式Attention(Q, K, V) softmax(Q·Kᵀ / √d)·VQ是查询向量K是键向量V是值向量。公式做的事情简单说就是拿当前的查询Q跟所有位置的K算相似度得到注意力权重再用这个权重去加权求和所有位置的V。多头注意力MHAMulti-Head Attention就是把这个过程重复H次每次都使用独立的Q、K、V线性投影。之所以要“多头”是因为单头注意力只能建模一种“查询-键”的匹配模式而一个句子里的依赖关系是多种多样的。有的头可能负责捕获代词指代有的头负责捕获语法依赖有的头负责捕获位置关系。每个头相当于一个“专家”从自己的角度观察序列最后把多个专家的报告拼在一起。这里有个容易被忽略的细节MHA里的每个头K和V都是独立的。假设有8个头每个头的维度是64那么同一个token在某一层里会对应8份不同的K向量和8份不同的V向量。这些K、V分别在不同子空间里编码信息互不干扰。训练阶段这种做法没有问题它给了模型极大的表达空间。但推理阶段的问题来了生成式大模型是自回归的一个token一个token往外蹦。每生成一个token当前token的Q要跟前面所有历史token的K做点积然后加权求和V。如果前面有1024个token这1024个token的K和V就要反复参与计算。而问题是这些历史token的K和V早在它们第一次出现时就已经算过了——后面再算一遍结果一模一样。1.2 MQA把KV的多头压缩成一个头MQAMulti-Query Attention的核心思路非常直接既然KV算一遍就够了那不如让所有Query头共享同一份K和V。K和V不再分H个头而是只保留一个共享头。这个名字里的“Multi-Query”指的是Q仍然多头但K、V变成单头。论文出处是Google 2019年的《Fast Transformer Decoding: One Write-Head is All You Need》。当时提出MQA目标很明确加速增量解码。因为解码阶段每生成一个token都要读取整段历史KVKV越小访存越少速度越快。用具体维度说话假设MHA配置是8个头、每头维度64。处理一个token时K的shape是[8, 64]V的shape也是[8, 64]。MQA把KV头压缩成1个后K的shape变成[1, 64]V也变成[1, 64]。8倍的压缩率。MQA带来的直接好处有两个。第一KV Cache体积直接除以H。比如一个7B模型的KV Cache在某些配置下需要2GB用MQA直接降到250MB左右。第二解码阶段每个token需要搬运的KV数据量同步变成原来的1/8访存压力大幅降低。代价也很明显KV被压成单头相当于所有Query头都必须在一个共享子空间里找信息。这就像一个团队所有人共用同一个文秘文秘再专业也扛不住同时服务8个需求迥异的业务线。实验也证实在中小规模模型上MQA相对MHA有质量损失尤其影响翻译、摘要这类对细粒度语义敏感的生成任务。MQA后来没有被大规模采用并不是因为它不好而是因为它太极端。它把KV压到极限换来了速度和显存却丢了一部分表达能力。这个“1”和“H”之间显然存在一个中间值——这就是GQA出现的动机。2. 从KV Cache的内存峰值反推GQA的分组妥协2.1 为什么单头KV会“塌方”表达能力临界点MQA质量下降的根因可以归结为KV角色的信息瓶颈。在多头注意力里每个头不只是“Query不同”它的K和V投影矩阵也不同这意味着每个头看到的是完全不同的特征子空间。多个头可以并行关注不同的关系模式有的头关注近距离的局部依赖有的头关注跨越很长距离的核心实体有的头负责编码句法结构。一旦KV变成单头所有Query头只能去查询同一个子空间。这个共享子空间必须同时满足所有Query头的查询需求很容易出现“众口难调”的局面。极端情况下模型为了让共享KV尽量兼顾所有人会选择“平均化”的表征——什么都包含一点但什么都不精确。对特别依赖细粒度特征的生成任务来说这点损失会被放大。但MQA的存在也揭示了一个重要事实注意力头之间存在大量冗余。很多研究通过可视化注意力头分布发现不少头学到的模式高度相似。既然有冗余就说明KV头数量不需要跟Query头一样多只是不能激进到只剩下1个。2.2 GQA按组分摊KV用可调节旋钮做折中GQAGrouped-Query Attention在2023年由Google的T5团队提出论文名是《GQA: Training Generalized Multi-Query Transformer with Multi-Query Linear Projections》。它的做法是把Query头分成G组每一组共享一份K和V。说直白点不是所有头共用一个KV而是组内共享组间独立。GQA的数学形式可以统一概括三个机制当G 1时所有Query头共享一份KV退化为MQA当G H时每个Query头独立一份KV退化为MHA当1 G H时就是真正的GQA。举个例子某模型有32个Query头GQA配置为8组每组4个Query头。KV头数量就是8个KV Cache体积降为MHA的1/4。Mistral 7B用的就是这类配置32个Query头、8个KV头。GQA的原理依据是注意力头的功能存在“聚类”现象。一部分头关注的是相同的语法模式一部分头关注的是相近的位置偏差。既然这些头做的事情高度相似它们完全可以共享底层KV信息只是在最终输出的线性拼接时保持独立。用一组KV服务多个靠近的Query头不会损失太多信息。论文里还有一个非常实用的工程技巧将标准MHA模型转换为GQA时不必从头训练。把MHA中同一个分组内对应的KV权重做平均得到GQA模型的初始KV权重再继续预训练或微调一小段时间就能恢复绝大部分效果。这也是很多开源模型在从实验版迭代到正式版时敢于直接换GQA架构的原因——它可以低成本复用已有权重。2.3 主流模型里的GQA配置我把主流开源模型的注意力配置拉了一张表读者可以直观感受一下GQA的地位。模型Query头数KV头数机制Llama 2 7B3232MHALlama 2 70B648GQA (8:1)Llama 3 8B328GQA (4:1)Mistral 7B328GQA (4:1)Qwen2 7B284GQA (7:1)Gemma 7B1616MHAFalcon 40B1288MQA为什么近两年的模型“默认”GQA核心驱动力不是效果而是推理服务方的成本账。大模型部署之后运营方最关心的指标是并发吞吐一块卡同时服务多少路请求。这个数字直接受KV Cache大小影响。KV Cache越小同样显存能囤的并发请求就越多单次请求成本就越低。GQA在表达能力损失可控的前提下把KV Cache压到了MHA的1/4甚至1/8成了商业量产模型的理性选择。这是一个典型的“推理优化倒逼架构设计”案例。模型结构不再只服务于训练效果还要为推理成本让路。这也是为什么你在看Llama 3、Mistral、Qwen等新模型的技术报告时几乎找不到还用纯MHA的大规模主力模型。3. KV Cache的“值班表”逻辑缓存什么、怎么算内存、何时释放3.1 缓存什么、为什么能缓存KV Cache按字面意思就是缓存K矩阵和V矩阵的缓存区。它在推理过程中被写入、反复读取、持续增长。为什么可以缓存关键点在于Self-Attention的计算特性某个token的K向量和V向量只取决于该token的输入表示和当前层的权重矩阵与后续出现的其他token无关。当前token不会“看到”后面的token再去改变自己的K和V它只需要在计算注意力权重时去跟其他token的K做匹配。打个比方KV Cache就像考试时用的草稿纸。你做完第一问把过程和结果贴在草稿纸角落做第二问、第三问时直接引用不需要每道题都从第一问重新推一遍。自回归生成时每一步新产生的token只需要用它自己的Q去查询已经写在草稿纸上的历史K再拿历史V做加权然后把自己这一轮的K和V追加到草稿纸上留给后面的token用。缓存的对象明确每一层Transformer的注意力计算都要缓存一组K和V。一个7B模型通常是32层每一层都有自己的KV Cache。所以KV Cache的总量是各层之和不是一层。3.2 一块KV Cache到底占多少显存KV Cache的内存计算有一个标准公式KV Cache内存 2K和V两套 × 层数 × 序列长度 × KV头数 × 每头维度 × 精度字节数用Llama 2 7B的实际参数代入看看。Llama 2 7B是32层、KV头4个实际上Llama 2 7B官方用的还是MHA这里为了演示GQA配置下的效果改成4个KV头为例、每头维度128、序列长度4096、FP16精度占2字节2 × 32 × 4096 × 4 × 128 × 2 268,435,456 字节 ≈ 256MB这是单条请求、4096上下文的KV Cache大小。对比一下如果是MHAKV头数32算出来就是约2GB。同一个模型仅仅因为KV头数从32砍到4KV Cache缩水了8倍。再极端一点上下文拉到32KMHA2 × 32 × 32768 × 32 × 128 × 2 ≈ 16GBGQA4个KV头2 × 32 × 32768 × 4 × 128 × 2 ≈ 2GB这就解释了一个现象为什么长上下文模型几乎标配GQA。因为纯MHA的模型把上下文拉长到32K后KV Cache能吃掉十几GB显存单卡根本跑不动。而GQA还能把KV控制在2GB级别预留了梯度空间。3.3 Prefill与Decode一写多读的工作节奏LLM推理分为两个阶段Prefill预填充和Decode解码。Prefill阶段处理的是用户输入的prompt。比如用户一次输入512个token这个阶段要把这512个token全部并行算一遍生成每一层每个token的K和V写入KV Cache。这个阶段是计算密集型的因为512个token可以同时过矩阵乘法GPU的算力利用率很高速度也快。KV Cache在这个阶段被“写满”了第一段。Decode阶段是逐token生成。每生成一个新token模型跑一遍前向传播但这个token的Q要跟KV Cache里所有历史的K做注意力然后加权V。所以每一层都要读取整段KV Cache但新增写入只有一个token的K和V。这个阶段是访存密集型的——读一大片写一点点。两个阶段的工作节奏完全不同。Prefill是“一口气写一篇论文”Decode是“每写一个字翻一遍所有参考文献”。这就是为什么开头提到的现象会出现大模型“读题”时是前缀并行计算几百毫秒搞定真正出字时每步都依赖全量历史KV快不起来。3.4 工程侧的KV Cache优化方向既然KV Cache是大模型推理的内存大头和速度瓶颈工程界自然不会坐视不管。目前主流的优化方向有几个PagedAttentionvLLM的核心技术。KV Cache不再一次性预分配最大长度的连续内存而是按固定块大小比如16个token一页动态分配。这样既避免了预分配浪费也能把物理上不连续的内存块串起来大幅减少内存碎片。实际部署中PagedAttention能把显存利用率从40%拉到90%以上。KV量化对缓存中的K和V做INT8、FP8甚至INT4量化减少字节数。KV量化跟权重量化不同KV是推理过程中动态产生的量化需要考虑值域分布随时间变化的问题实现起来比静态量化麻烦。但收益直接KV从FP16变INT8显存和带宽同时减半。滑动窗口注意力Mistral用的方案只缓存最近N个token的KV更早的KV直接丢。适合流式生成场景因为当前token对很久之前的依赖通常不明显。代价是超过窗口长度的远距离依赖会丢失。Token DroppingH2O等方案的思想是每步只保留注意力权重比较高的“重要token”的KV权重低的KV删掉。相当于动态压缩历史能在长上下文下保持较高精度。这些优化手段都指向同一个目标KV Cache太大、读取太频繁。MQA和GQA在结构上做减法这些方法在工程上做压缩两者可以叠加使用实际部署时经常一起上。4. 显存带宽是主角为什么推理瓶颈不在计算而在搬运4.1 Decode阶段的Attention为什么这么“吃带宽”很多人的直觉是大模型推理应该很吃算力毕竟矩阵乘那么大。但实际decode阶段的Attention计算量并不大。看一个输出token的计算过程。当前token的Q是一个维度为d的向量跟Cache里L个历史token的KL×d矩阵做点积得到L个注意力分数。再拿这L个分数去加权L个V向量也是L×d。整个过程的浮点运算量大约是 2 × L × d规模不大。但搬运的数据量是L × d的全量KV。假设d是128、L是4096一份K是4096 × 128的矩阵占1MBFP16一份V也占1MB所有层加在一起这个搬运量被放到几GB的级别。问题就出在“计算少、搬运多”上。GPU计算单元速度快但从显存把数据搬到计算单元的速度带宽是有限的。当搬运一个数据的时间远大于计算这个数据的时间整个任务就被带宽卡住了。技术上叫memory-bound计算单元的利用率很低GPU在大量空转等数据。类比一下计算单元像一位世界顶级大厨显存像冰箱。大厨做一道菜只需要10秒但从冰箱把菜拿出来要走3分钟。大厨再快也没用时间全花在“取菜”上了。4.2 用真实带宽算一笔账以A100 80G为例它的HBM显存带宽约2TB/s。假设某个模型的KV Cache是1GB每生成一个tokenAttention部分至少要读1GB的KV数据有些实现还要读两遍加上模型权重、激活值、中间结果实际搬运量更高。算一下理论极限2TB/s ÷ 1GB ≈ 2000 tokens/s。听着很快但这是理想上限实际单用户生成速度远低于这个值。因为第一模型权重也要从显存读取。跑一个7B模型每生成一个token光把权重从显存过一遍就是14GBFP167B×2字节。这个重量级搬运远远超过KV Cache。第二KV Cache读取可能发生多次。不同层的KV存储位置不同读取路径有开销。第三Attention计算中softmax、归一化等操作会让计算单元有等待时间。综合下来真实部署中单卡7B模型能跑出50-100 tokens/s已经算不错。如果把KV从MHA换成MQA或GQAKV Cache从2GB降到250MB带宽压力立刻小很多生成速度能上一个台阶。4.3 MHA/MQA/GQA在读带宽上的差距这里说一个容易混淆的点MQA的KV Cache大小不一定比GQA小。KV Cache大小取决于num_key_value_headsKV头数。MQA是KV头数等于1GQA是KV头数等于GG1。比如Llama 3 8B的配置是32个Query头、8个KV头GQA。它的KV Cache跟一个32Query头、1个KV头的MQA模型相比是后者几乎没有表达能力但两者的KV Cache都是250MB级别假设其他配置相同。所以准确的说法是MQA和GQA都能把KV Cache压到很小差距在于KV头数。KV头数越少KV Cache越小带宽压力也越小。GQA的优势不在于比MQA更省而在于同样省内存的代价下保留了多组KV表达能力更好、训练更稳定。在实际推理中KV Cache的读取带宽消耗与KV头数线性相关。KV头数从32降到8Attention部分的带宽消耗直接降到1/4。这也是为什么在做推理性能调优时看模型的num_key_value_heads是一个很重要的参考指标。5. 实际工程中的选型参考与显存估算经验5.1 从config一眼识别MHA/MQA/GQA拿到一个开源模型想判断它用的是MHA、MQA还是GQA不用翻论文直接看config.json里的两个字段num_attention_heads和num_key_value_heads。{ num_attention_heads: 32, num_key_value_heads: 8, num_hidden_layers: 32, hidden_size: 4096 }判断规则很简单num_key_value_heads num_attention_heads → MHAnum_key_value_heads 1 → MQA1 num_key_value_heads num_attention_heads → GQAHuggingFace的transformers库在加载模型时会根据num_key_value_heads自动决定用哪种注意力实现。如果你自己写推理代码这两个字段是计算KV Cache大小的关键输入。5.2 完整显存估算与并发度推算部署一个模型之前最好先估算显存占用。总显存 模型权重 激活值 KV Cache 服务框架开销。模型权重最简单参数量 × 精度字节数。7B模型FP16就是14GBINT4量化后大约3.5-4GB。激活值中间计算结果比较难估跟batch size和序列长度有关一般粗略按几GB估算。服务框架vLLM、TensorRT-LLM等自身也有显存开销。KV Cache按前面的公式算。举个例子一台16G显存的消费级显卡比如RTX 4090部署Qwen2-7B。4bit量化权重大约5GB激活值预留2GBKV Cache按上下文8K、GQA配置假设4个KV头算2 × 28层 × 8192 × 4 × 128 × 2字节 ≈ 469MB这里的28层是Qwen2-7B的实际层数。算下来KV Cache不到0.5GB这给并发留了很大空间。但如果换成同规模MHA模型KV头数变成28个KV Cache直接飙到3.3GB。同样的显卡配置差距就这么大。实际部署时vLLM等框架会按最大并发数预先分配KV Cache空间。估算单卡能跑多少并发公式是可并发数 ≈ (总显存 - 权重显存 - 激活显存 - 框架开销) ÷ 单条序列KV Cache保守估算16G显存权重5G激活2G框架1G剩余8G。单条序列KV Cache按1K上下文算约60MB那么并发可以开到100以上。在本地玩完全够了。但如果你把max_length拉到32K、并行跑多条promptKV Cache暴涨OOM就在所难免。5.3 踩坑记录与经典误区几个我实际踩过、也看别人反复踩的坑误区一KV Cache要参与反向传播。对微调和训练来说KV Cache确实需要保存中间值用于梯度计算。但纯推理场景KV Cache只做前向计算不需要存梯度。有些新手自己手写推理脚本时把训练逻辑搬过来导致显存翻倍浪费。误区二max_length设多大都行。很多人部署时只看模型权重大小忽略了max_length对KV Cache的影响。把max_length从4K调到32KKV Cache涨8倍明显准备不足就直接OOM。实际部署时应按业务场景合理设置上下文长度不要盲目拉满。误区三GQA一定比MHA效果差。在足够大的模型规模下GQA的下游任务表现跟MHA很接近。Meta在Llama 2 70B上专门做过对比GQA在多数指标上跟MHA差距在1个点以内但在推理效率和长上下文表现上优势明显。中小模型上可能差距更大一些但“GQA劣化”是一个过于简化甚至错误的结论。坑微调时随意改num_key_value_heads。如果直接用开源预训练权重做微调注意不要改num_key_value_heads这个参数。改了之后结构就对不上预训练权重了加载时报错。如果非要用GQA的模型结构应该去找本来就是GQA的预训练模型而不是自己魔改结构。还有一个实操建议如果你要在消费级显卡上本地部署大模型优先选GQA架构的小模型7B-14B级别把上下文长度控制在8K以内。这样KV Cache通常不会成为瓶颈剩下的优化精力可以放在权重量化和服务框架选择上。我自己做推理优化时最大的体会是KV Cache不是一道背定义就能过的八股题而是理解大模型推理成本最重要的一把尺子。很多同学背熟了MHA、MQA、GQA的定义部署时却不知道max_length该调多小也不知道为什么vLLM会疯狂吃显存。如果你真想把这套东西吃透不一定要从零撸代码可以先拿一个GQA小模型在本地跑起来开着显存监控工具不断拉长上下文亲眼看KV Cache是怎么一步步把显存吃掉的。这一趟下来比看十篇原理文章都管用。
返回列表