08 · Multi-Query Attention (MQA)

注意力 Shazeer 2019
Noam Shazeer, Google · arXiv:1911.02150

迭代了什么

上一代的问题:多头注意力(MHA)推理时 KV 缓存巨大,每个头各自存 K 和 V(h 份)。

这代改了什么:所有头共享同一组 Key 和 Value,只保留 Query 的多头性。KV 缓存降为 1/h。

效果:推理加速 ~12 倍(3.8μs vs 46μs 每步)。PaLM 采用。困惑度轻微上升(1.424→1.439)。

一、解决的问题

Transformer 自回归推理时,多头注意力需要反复加载完整的 KV 缓存。内存带宽成为推理瓶颈。

二、核心思想

所有注意力头共享同一组 Key 和 Value,只有 Query 保持多头。KV 缓存从 O(h·n·dh) 降为 O(n·dh)。

三、公式

MHA: K, V ∈ [b, h, m, d]
MQA: K, V ∈ [b, m, d] (去掉 h 维度)

四、关键权衡

推理加速约 12 倍,质量轻微下降(困惑度 1.424 → 1.439)。PaLM 使用 MQA。