1. 从“推理卡顿”说起为什么我们需要关注KV缓存与注意力机制最近在折腾一个基于Transformer的文本生成项目模型不大也就几十亿参数。在本地用单卡跑推理测试时我发现一个挺有意思的现象生成前几个token时速度飞快但越往后生成速度就越慢甚至能感觉到明显的卡顿。这显然不符合直觉——模型参数是固定的计算量难道不是每个token都差不多吗起初我怀疑是显存带宽瓶颈或者Python的GIL锁问题但一通Profiling性能剖析下来发现瓶颈并不在数据传输或Python解释器上。真正吃掉大部分时间的是一个叫做“注意力计算”Attention的环节而且时间消耗随着已生成序列的长度平方级增长。这就像你每说一个新词都要把之前说过的所有词从头到尾再回忆、比对一遍话越长回忆的过程就越吃力。这个问题的核心就是Transformer解码器在自回归生成Autoregressive Generation时的固有特性。为了生成下一个token模型需要基于之前所有已生成的token来计算注意力。而存储这些历史token的Key和Value向量就是KV缓存KV Cache。没有它每次生成都需要重新计算整个历史序列的注意力开销无法承受但简单粗暴地缓存所有KV又会带来巨大的内存开销和访存压力尤其是在生成长文本时。更棘手的是标准的**多头注意力Multi-Head Attention, MHA**机制中每个注意力头都有自己独立的Key和Value投影这导致KV缓存的大小与注意力头数成正比。当模型规模增长到千亿参数、拥有上百个注意力头时KV缓存的内存占用会成为一个非常恐怖的负担直接限制了我们能生成的序列长度也拖慢了推理速度。于是工程师们开始寻找既能保持模型表达能力又能显著压缩KV缓存大小、提升推理效率的注意力变体。分组查询注意力Grouped-Query Attention, GQA就是在这样的背景下从研究论文走向工程实践的一个关键优化。它本质上是一种在多头注意力MHA与多查询注意力Multi-Query Attention, MQA之间的优雅折中。简单来说你可以这样理解这三者的关系MHA多头注意力追求极致的模型表达能力每个头都有独立的K、V但KV缓存开销大。MQA多查询注意力追求极致的推理效率所有头共享同一份K、V缓存开销最小但可能牺牲过多模型能力。GQA分组查询注意力一种灵活的“分组套餐”。将多个头分成一组组内共享同一份K、V。比如8个头分成2组那就有2份独立的K、V。它在缓存大小和模型能力之间取得了更好的平衡。接下来我们就深入这个“推理卡顿”问题的核心拆解KV缓存的工作原理、它带来的挑战并重点剖析GQA是如何通过改变注意力头的组织方式来显著缓解内存和带宽压力从而让大模型推理变得更流畅、更经济的。无论你是正在部署模型的服务端工程师还是对Transformer底层机制感兴趣的研究者理解这些概念都至关重要。2. KV缓存Transformer自回归推理的“记忆体”与性能双刃剑要理解GQA的价值必须先彻底搞懂KV缓存是什么以及它为什么如此重要又如此“麻烦”。2.1 自回归生成与重复计算的陷阱Transformer解码器如GPT系列生成文本的方式是自回归的给定一个初始输入提示词模型输出第一个token然后将这个token拼接到输入中作为新的输入再输出下一个token如此循环往复。在标准的Transformer注意力机制中第t步要计算当前查询向量Q_t与之前所有步的键向量K_{1:t}和值向量V_{1:t}的注意力。如果我们不缓存任何东西那么在生成第t个token时就需要将前t-1个token的输入重新通过模型的前馈层和注意力投影层计算出它们对应的K_{1:t-1}和V_{1:t-1}。这意味着大量重复计算计算复杂度是O(n^2)完全不可行。2.2 KV缓存的工作原理空间换时间KV缓存的核心思想非常直接既然每一步生成的K和V只依赖于当前的输入token并且在后续生成中不会再改变那么我们就可以把它们存储缓存起来供后续步骤直接使用。具体流程如下预填充阶段Prefill处理用户输入的提示词Prompt。对于提示词中的每一个token模型正常计算其对应的Q,K,V。其中K和V会被存储到缓存区中。这个阶段是“写入”缓存。生成阶段Decoding开始生成第一个新token。对于当前要生成的新位置比如第一个新token的位置模型计算其Q_new。从缓存中读取之前所有步骤包括提示词和已生成token存储的K_cache和V_cache。将Q_new与K_cache进行注意力计算得到权重再作用于V_cache得到当前步的上下文向量。模型输出该位置的token。关键一步将这个新生成的token输入模型计算出它对应的K_new和V_new然后将它们追加到缓存中。重复第2步直到生成结束。这个过程将每一步的计算复杂度从O(n^2)降低到了O(n)主要是注意力计算中的矩阵乘是一种典型的以空间内存换取时间计算的策略。2.3 KV缓存的内存开销一个具体的计算示例KV缓存的开销是实实在在的。我们以一个流行的开源模型LLaMA-2 70B的参数为例进行估算隐藏层维度hidden_size: 8192注意力头数num_heads: 64每头维度head_dim: 8192 / 64 128层数num_layers: 80精度dtype: 通常推理使用 float16 (2字节)对于标准的多头注意力MHA每一层、每一个注意力头都有自己独立的K和V投影权重。在推理时对于序列中每一个token我们需要为每一层缓存它的K和V向量。每个token每层的KV缓存大小 num_heads * head_dim * 2(K和V) *bytes_per_parameter代入数值64 * 128 * 2 * 2字节 32768 字节 32 KB对于80层32 KB * 80 2560 KB ≈ 2.5 MB这意味着生成或处理一个token仅KV缓存就需要约2.5 MB的存储空间。这还只是一个token。如果我们想要生成一个长度为2048的序列KV缓存总大小 2.5 MB/Token * 2048 Tokens ≈ 5120 MB 5 GB这5GB是额外的显存占用不包括模型参数本身约140GB FP16和激活值等。对于一张40GB显存的A100显卡这5GB的缓存就吃掉了八分之一的显存严重限制了批量大小Batch Size和可处理的序列长度。在服务端场景高并发意味着需要同时处理多个请求每个请求都有自己的KV缓存显存压力会成倍增加。注意上述计算是近似值实际实现中可能会因框架优化如连续存储、内存对齐而略有不同但数量级是准确的。这个开销清晰地表明MHA的KV缓存是长序列推理的主要瓶颈之一。3. 从MHA到MQA注意力机制的效率演进为了应对KV缓存的开销问题学术界和工业界首先提出了一种激进的方案多查询注意力Multi-Query Attention, MQA。理解MQA是理解GQA的基础。3.1 回顾标准多头注意力MHA的冗余在MHA中假设有h个头对于输入序列X经过线性投影后Q X * W_q- 形状:[batch, seq_len, h * d_k]通常拆分为[batch, seq_len, h, d_k]K X * W_k- 形状:[batch, seq_len, h * d_k]拆分为[batch, seq_len, h, d_k]V X * W_v- 形状:[batch, seq_len, h * d_v]拆分为[batch, seq_len, h, d_v]每个头i使用自己的Q_i,K_i,V_i计算注意力。这里的核心在于K和V的投影权重W_k和W_v是每个头独立的。这带来了表达能力的灵活性每个头可以关注不同的特征但也导致了KV缓存与头数h成正比。3.2 多查询注意力MQA的极致压缩MQA提出了一个大胆的简化让所有的注意力头共享同一套Key和Value投影。Q X * W_q- 形状:[batch, seq_len, h * d_k]拆分为[batch, seq_len, h, d_k](不变)K X * W_k- 形状:[batch, seq_len, d_k](关键变化没有头维度h)V X * W_v- 形状:[batch, seq_len, d_v](关键变化没有头维度h)在计算注意力时对于每一个头i其Q_i会与这个共享的K进行计算输出的上下文向量由这个共享的V加权得到。带来的好处是革命性的KV缓存大小急剧减少缓存不再与头数h相关。每个token每层的KV缓存大小从h * d_k * 2降为d_k * 2。以上述LLaMA-2 70B为例缓存大小直接减少为原来的1/64。单个token每层缓存从32KB降到0.5KB80层仅需40KB。生成2048个token也只需约80MB缓存相比MHA的5GB减少了98%以上内存带宽压力骤降在生成阶段每一步都需要从显存中读取整个KV缓存。MQA使得每次读取的数据量减少了h倍极大地缓解了内存带宽瓶颈从而提升了推理速度。计算图简化某些矩阵运算的维度降低带来微小的计算加速。3.3 MQA的潜在代价表达能力下降然而天下没有免费的午餐。MQA的激进压缩可能带来模型表达能力的下降。表征多样性受限在MHA中不同的头可以学习到关注输入序列中不同方面或不同位置的信息。共享K和V意味着所有头都基于同一套“记忆”进行检索可能限制了模型捕捉复杂模式和关系的能力。训练稳定性与性能一些研究发现直接从零开始训练一个MQA架构的模型有时在最终性能上会略逊于同参数量的MHA模型。尤其是在需要精细理解上下文或进行复杂推理的任务上。MQA是一种“效率优先”的架构它在许多场景下特别是对话、续写等常见任务表现足够好且收益巨大。但当我们对模型能力有极致要求时就需要一个更平衡的方案。4. 分组查询注意力GQA在效率与能力间寻找黄金分割点GQA的设计哲学非常直观既然MHA太“胖”缓存大MQA又可能太“瘦”能力可能受损那我们为什么不取一个中间状态呢4.1 GQA的核心思想分组共享GQA将原始的h个注意力头分成g个组group。每个组内包含h/g个头假设h能被g整除。每个组拥有自己独立的一套Key和Value投影权重。同一个组内的所有头共享这套KV投影。用公式和形状来表示会更清晰设头数h 8 组数g 2 则每组有4个头。Q X * W_q- 形状:[batch, seq_len, h * d_k]- 视图为[batch, seq_len, g, h/g, d_k]K X * W_k- 形状:[batch, seq_len, g * d_k](注意维度是 g * d_k 不是 h * d_k)V X * W_v- 形状:[batch, seq_len, g * d_v]在计算时属于第j组的那些头会使用第j组的K_j和V_j进行计算。这带来了灵活的配置空间当g h时GQA 退化为 MHA每组1个头各自独立。当g 1时GQA 退化为 MQA所有头为一组完全共享。当1 g h时就是典型的GQA。例如h64,g8则每组8个头共享KV。4.2 GQA带来的收益分析KV缓存的有效压缩缓存大小与组数g成正比而不是头数h。压缩比为g / h。例如对于64头的模型采用8组GQAKV缓存大小就降为MHA的1/8。这依然是一个巨大的节省同时保留了分组内的表征多样性。内存带宽压力成比例降低与缓存减少同步每一步读取KV缓存的数据量也降为原来的g/h显著提升解码速度。保持模型能力通过分组模型仍然保留了多组不同的“记忆视角”。不同组可以学习关注输入的不同子空间或特征理论上比单一的MQA具有更强的表达能力。实践也证明通过恰当的训练包括从MHA模型进行蒸馏GQA模型可以在几乎不损失精度的情况下获得接近MQA的推理效率。4.3 GQA的训练策略从MHA进行上采样与蒸馏一个常见的问题是如何得到一个GQA模型有两种主要方式从头训练直接使用GQA架构定义模型并用大量数据从头训练。这需要大量的算力和数据但能确保模型从头学习分组共享的表示。从预训练MHA模型转换更流行这是目前更实用的方法。以一个训练好的MHA模型如LLaMA-2为起点通过“上采样”和蒸馏来获得GQA模型。上采样Upsampling对于MHA模型的每一层我们有h套独立的W_k和W_v权重。要转换成g组的GQA我们需要将这h套权重“融合”成g套。一种简单有效的方法是平均池化将h个头分成g组将每组内所有头的W_k和W_v参数取平均作为新组的权重。蒸馏微调使用上采样得到的GQA模型作为初始化在少量数据上甚至可以是原训练数据的一个子集进行短暂的继续预训练或指令微调。让模型适应新的分组注意力机制恢复可能因权重平均而损失的少量性能。这个过程计算成本相对较低。许多最新的开源和闭源大模型都采用了GQA。例如Google的Gemma模型家族就明确使用了GQA。Meta的LLaMA-2是MHA但社区有其GQA变体。Mistral AI的模型也采用了类似GQA的架构。这已经成为大模型推理部署的事实标准之一。5. 工程实践在推理框架中实现与管理GQA与KV缓存理解了原理我们来看看在真实的推理引擎如vLLM, TensorRT-LLM, Hugging Face Transformers中这些东西是如何落地的。5.1 KV缓存的内存布局与高效管理KV缓存的管理是推理引擎的核心优化点。简单地在内存中开辟一个大数组来存储所有token的KV是低效的。现代推理框架会采用更精细的策略连续内存与分块存储为了优化内存访问模式提高缓存命中率KV缓存通常被组织成连续的内存块。例如vLLM提出了PagedAttention思想将KV缓存划分为固定大小的“块”blocks类似于操作系统中的内存分页。每个请求的KV缓存可以分散在不连续的物理块中但通过逻辑块表来管理。这极大地减少了内存碎片允许更灵活的动态序列长度支持。内存复用对于批处理batch中的多个请求如果它们的提示词有公共前缀common prefix这部分前缀的KV缓存可以被多个请求共享避免重复存储。缓存逐出与压缩在内存紧张时可以结合注意力分数等信息尝试丢弃或压缩一些相对不重要的历史KV尽管这属于更高级的优化可能影响生成质量。5.2 集成GQA计算图的改写与内核优化在支持GQA的推理框架中需要实现对应的计算内核kernel。形状变换框架需要识别模型的注意力层是GQA配置通过配置参数如num_attention_heads和num_key_value_heads来指定。在计算时会将Q张量 reshape 为[batch, seq_len, num_attention_heads, head_dim]而将K和Vreshape 为[batch, seq_len, num_key_value_heads, head_dim]。这里的num_key_value_heads就是GQA中的组数g。广播计算在计算Q K^T时由于K的头维度g小于Q的头维度h需要将K沿着头维度进行广播broadcast以便与每个Q头进行计算。高效实现这一点需要定制的CUDA内核或巧妙使用现有张量运算。融合内核为了极致性能像FlashAttention这样的优化注意力内核也需要推出支持GQA/MQA的版本。FlashAttention通过算子融合和精细的GPU内存管理来加速注意力计算其对GQA的支持能带来端到端的显著加速。5.3 实操示例在Hugging Face Transformers中使用GQA模型以使用一个假设的、支持GQA的模型为例例如mistralai/Mistral-7B-v0.1 注意实际Mistral-7B是使用滑动窗口注意力SWA但这里我们用其配置概念演示GQA。from transformers import AutoTokenizer, AutoModelForCausalLM import torch # 加载模型和分词器 model_name mistralai/Mistral-7B-v0.1 # 此处仅为示例实际模型架构可能不同 tokenizer AutoTokenizer.from_pretrained(model_name) model AutoModelForCausalLM.from_pretrained(model_name, torch_dtypetorch.float16, device_mapauto) # 查看模型配置中的关键参数 print(model.config) # 通常会看到类似 # num_attention_heads32 # num_key_value_heads8 # 这个就是GQA的组数g意味着是32头分8组。 # hidden_size4096 # head_dim hidden_size / num_attention_heads 128 # 准备输入 prompt 请解释一下人工智能。 inputs tokenizer(prompt, return_tensorspt).to(model.device) # 生成时Transformers库会自动处理KV缓存和GQA计算 with torch.no_grad(): outputs model.generate(**inputs, max_new_tokens100, do_sampleTrue, temperature0.7) print(tokenizer.decode(outputs[0], skip_special_tokensTrue))在底层transformers库的generate函数会自动管理一个past_key_values元组它就是KV缓存。对于GQA模型这个缓存中每个层存储的K和V张量的形状其头维度会是num_key_value_heads而不是num_attention_heads。5.4 性能对比与选型建议在选择MHA、MQA还是GQA时需要权衡特性多头注意力 (MHA)多查询注意力 (MQA)分组查询注意力 (GQA)KV缓存开销大 (O(h))极小 (O(1))中等 (O(g))内存带宽压力高极低低理论表达能力最高可能较低较高 (接近MHA)训练难度标准可能需更多数据/技巧可从MHA蒸馏相对容易适用场景研究、对精度要求极高的场景极度追求吞吐和低延迟的推理场景生产环境推理的平衡之选个人建议对于全新的模型训练如果推理效率是重要考量可以优先考虑采用GQA作为默认架构。对于基于现有MHA模型进行部署如果面临显存或速度瓶颈强烈建议尝试将其转换为GQA模型通过上采样蒸馏。这是一个成本相对较低、收益显著的优化手段。只有在推理资源极其紧张如边缘设备且对精度损失有一定容忍度时才考虑使用MQA。6. 超越GQAKV缓存优化的其他前沿思路GQA主要解决了KV缓存大小和带宽的问题。但在长文本生成如处理数万token的上下文场景下即使使用GQA缓存总量依然会线性增长。社区还在探索更激进的优化方向6.1 选择性缓存与动态稀疏化核心思想是不是所有历史token的KV都同等重要。我们可以选择性地缓存那些“重要”的KV。基于注意力分数的淘汰定期检查历史token的注意力分数均值淘汰那些长期不被关注的token的KV。基于信息熵的压缩尝试将多个不重要的token的KV向量合并或量化用更少的空间存储近似信息。滑动窗口缓存像Mistral的滑动窗口注意力Sliding Window Attention, SWA一样只缓存最近W个token的KV。这严格限制了缓存大小但模型必须被专门训练以适应这种“有限记忆”。这些方法属于“有损压缩”需要在效率和质量之间做精细的权衡目前大多处于研究阶段。6.2 量化与低精度存储这是目前最直接、应用最广泛的压缩手段。KV缓存量化将KV缓存从FP162字节量化为INT81字节甚至INT40.5字节。这可以直接将缓存大小减半或降至四分之一。挑战注意力计算Q K^T需要高精度点积来维持注意力权重的准确性。因此通常需要将低精度的K缓存反量化到较高精度如FP16后再进行计算或者使用混合精度策略。支持主流推理框架如vLLM、TensorRT-LLM都已支持KV缓存的INT8量化并能与FlashAttention等优化内核结合在几乎不损失精度的情况下获得显著的显存节省和速度提升。6.3 内存与计算的重新权衡MQA、GQA与FlashAttention的协同未来的优化趋势是多层次、协同的。架构层采用GQA/MQA减少缓存的基本单元数量。系统层使用类似PagedAttention的内存管理技术减少碎片提高利用率。数值精度层对KV缓存进行量化。计算内核层使用高度优化的、支持上述所有特性分组、量化、分页的融合注意力内核如FlashAttention的变种。例如一个理想的系统可能这样工作模型采用GQA架构KV缓存以INT4精度存储在以“块”为单位管理的虚拟内存中。当需要计算注意力时调度器将所需的缓存块加载到SRAM并在内核中动态反量化为FP16与FP16精度的Q进行基于FlashAttention算法的融合计算。7. 总结与个人踩坑心得回顾一下KV缓存是Transformer自回归推理的必需品但也是性能瓶颈。GQA通过让多个注意力头分组共享KV投影在几乎不损失模型能力的前提下将缓存大小和内存带宽压力降低了数倍乃至数十倍是目前大模型推理部署中一项至关重要的技术。从我自己的实践来看有几点深刻的体会第一Profiling性能剖析永远是第一步。不要凭感觉猜测瓶颈。用Nsight Systems、PyTorch Profiler等工具清晰地看到时间花在了Attention计算、内存拷贝还是别的什么地方。我最初就是靠Profiler才锁定KV缓存读取是拖慢长文本生成的元凶。第二量化是“性价比”最高的优化手段之一。在尝试更复杂的架构改动如转GQA之前不妨先试试对KV缓存进行INT8量化。很多框架提供开箱即用的支持通常只需要加几行配置代码就能获得立竿见影的显存节省而精度损失在大多数情况下微乎其微。第三从MHA转换到GQA时蒸馏数据的选择很重要。如果你正在将一个预训练的MHA模型转换为GQA用于蒸馏微调的数据不一定要很大但质量和代表性很关键。最好使用与原任务领域相关的数据或者包含各种复杂推理、长上下文理解样本的数据集这有助于模型更好地恢复分组共享后的表征能力。最后保持对底层原理的好奇。理解GQA、KV缓存、PagedAttention、FlashAttention这些概念不仅能帮助你在使用高层次API时做出正确选择更能在遇到诡异问题比如生成结果偶尔出错、长文本后质量下降时有深入排查的方向。大模型推理是一个系统工程每一个环节的优化最终累积起来就是巨大的成本差异和用户体验提升。