大模型逐词生成时,前面每个 Token 的 Key 和 Value 都要保存起来。新 Token 进来,GPU 再把这些缓存读出,供各个注意力头计算。序列越长、批量越大,读写的量越可观。多查询注意力的做法很直接,查询仍保留多个头,Key 和 Value 却只保留一套,所有头共用。
我核对了提出 MQA 的原始论文和实验表。论文关心的是增量解码中的内存带宽。多头注意力为每个头分别保存 Key 与 Value,MQA 删去了缓存里的头维度。论文的理论分析认为,与序列长度有关的那部分带宽压力,可以约按注意力头数缩小。
当年的译文实验快了很多
原论文用 WMT 2014 英德翻译做对照。基线是 6 层、8 个注意力头的编码器解码器 Transformer,约 2.11 亿参数。MQA 版本把前馈层加宽,让总参数与基线保持相同,因此表里的差别不是简单由参数量减少带来的。每个模型训练 10 万步,每批 128 个样本,输入和目标各为 256 Token,用 32 核 TPUv3 约训两小时。
在一块 8 核 TPUv2 上,训练速度几乎没变。基线每个输出 Token 折算 13.2 微秒,MQA 是 13.0 微秒。差距出现在逐步生成。批量为 1024,输入和输出各 128 Token 时,基线解码器每个 Token 折算 46 微秒,MQA 降到 3.8 微秒。Beam 4 搜索中,203 微秒降到 32 微秒。这些数字说明了带宽问题,却不能原样套到今天的 GPU 和推理框架上。
质量付出很小,但没有消失。开发集上,多头基线的 BLEU 为 26.7,MQA 为 26.5;开发集对数困惑度从 1.424 变为 1.439。WMT14 测试集用 Beam 4 时,MQA 得到 28.5 BLEU,基线为 28.4。另一组十亿词语言建模实验中,困惑度从 29.9 升到 30.2。两组结果都支持一个有边界的判断,MQA 能大幅减少解码读取,模型质量可能略降。
同一张表还比较了直接减少头数或 Key、Value 维度。这些版本也把前馈层加宽到相同参数量,开发集 BLEU 只有 25.8 到 26.2,低于 MQA 的 26.5。保留多个 Query 头,比整体削掉注意力头更能守住质量。
作者还将局部注意力与 MQA 放在一起。解码器每层只看当前位置和前 31 个位置时,普通多头版本每 Token 用 23 微秒,再换成 MQA 后是 3.3 微秒。这说明减少关注位置和删去 KV 的头维度可以同时使用,两者解决的内存读取不同。实验为固定形状而把缓存补到 128,刚开始生成的延迟会因这个实现选择被高估。
现在部署要看完整实现
MQA 与减少查询头不同。它还让多个查询头保留不同的查询投影,只把 Key 与 Value 合并。分组查询注意力则介于多头和 MQA 之间,用多组 Key 与 Value 换取质量与缓存的折中。
评估 MQA 时,我会先算每层、每个 Token 实际保存了多少 KV,再用目标批量和序列长度测吞吐、单请求延迟与业务任务分数。缓存缩小是结构事实,最后快多少,仍取决于算子、硬件、批处理和调度。