查询头保持不变,只减少键值头
一句话总结
区分 MHA、MQA、GQA 的旋钮只有一个——设置多少个 key·value 头。query 头的数量不去碰。所以减少的是 KV 缓存占用的内存,而注意力的乘法次数几乎不变。
为什么需要它
生成是一个令牌一个令牌地前进的。每生成一个新令牌,模型都要重新查看到目前为止所有位置的 key 和 value。所以不把这些 key 和 value 扔掉再重新生成,而是一直带着。这就是 KV 缓存。
问题在于,这个缓存会随着上下文变长、层数变多、同时处理的请求变多而原样变大。模型权重与上下文长度无关,是固定的,缓存却不是。把权重全部装载之后,剩下的空间被缓存占去,所以从某个点开始,缓存决定了并发处理数。
而且生成一个令牌时做的计算并不多。为了做这点少量的计算,却必须把整个缓存从内存里读过来。MQA 原论文作为增量解码的瓶颈所指的,正是这种回读。标题本身就是答案——要用的头一个就够了。
MQA 的决定由此而来。只设一个 key·value 头。query 头不论是八个还是六十四个都保持原样,它们全都查看同一份 K·V。缓存缩小到 query 头数分之一。
代价是有所失去。本来设计成让每个头去看不同地方的结构里,决定“要拿出什么”的一侧被合并成了一个。GQA 原论文在两者之间安排了一个位置。把 query 头分成若干组,每组各设一个 key·value 头。
三个名字是同一条轴上的三个位置
假设 query 头有 h 个,组有 g 个。这时 key·value 头的数量就是 g。
| 组数 g | 名称 | key·value 头 | 查看同一份 K·V 的 query 头 |
|---|---|---|---|
| g = h | MHA | h 个 | 1 个 |
| 1 < g < h | GQA | g 个 | h/g 个 |
| g = 1 | MQA | 1 个 | h 个 |
不需要把这三行分开去背。它们的区别在于把同一个旋钮转到了哪里。条件只有一个——h 必须能被 g 整除。否则各组的大小就会参差不齐,有的组合并三个头,有的组合并两个头。在实际的配置文件里违反这个条件,通常在加载模型时就会被拦下。
什么减少,什么不减少
这是本主题的核心,也是最容易搞错的地方。把式子列出来,一眼就能看明白。
放进缓存的实数个数是这样的。
KV 캐시 원소 수 = 2 × 층 수 × (키·값 헤드 수) × 길이 × 헤드 차원
前面的 2 是 K 和 V 两份。这个式子里没有 query 头的数量。因为 query 不会留在缓存里。所以把 key·value 头从八个减到一个,缓存恰好变成八分之一。
接下来数一数乘法次数。这是长度为 n、query 头为 h、头维度为 d 的一层通过一次的情形。
점수 Q·Kᵀ h × n × n × d ← g 가 없다
가중합 × V h × n × n × d ← g 가 없다
K 투영 n × d_model × (g × d)
V 투영 n × d_model × (g × d)
注意力主体里没有 g。原因一看实现就清楚。使用折叠起来的 key·value 头时,要按 query 头的数量重复展开,然后照常运行。不同值的种类减少了,但注意力运行的次数由 query 头的数量决定。减少的项只有 K·V 的两个投影,而它们在整体中并不占大头。
于是可以这样总结。GQA 减少的是内存。变快是因为读那部分内存读得少了,而不是因为乘法做得少了。如果分不清这一点,就无法回答“换成了 GQA,为什么预填充还是老样子”这个问题。
怎样合并各组
假设要把已经训练好的多头模型迁移成 GQA。要把一组里的四个 key·value 头合成一个,该怎么合成呢?
可以只留一个,扔掉其余的;可以取平均;也可以按缩小后的尺寸从头重新训练。GQA 原论文选的路是把组内的 key·value 头取平均作为初始值,再从这里多训练一点。重要的是,并不是只取平均就结束了——平均是出发点,不是答案。
下一个实验也用平均。不过不做训练。所以在本实验中出现的“输出变化了这么多”这个数字,并不表示质量会变差这么多,而是显示在同样的权重下,只把 key·value 合并起来,答案就会不同这个事实的值。这种区分在报告里也要写下来。
在现场相遇的样子
第一,配置文件里头数写成了两行。如今的模型配置里,query 头数和 key·value 头数是分开的。两个值相同就是 MHA,key·value 一侧是 1 就是 MQA,介于两者之间就是 GQA。不必去找三个名字,看两个数字的比值就行。
第二,会冒出“明明是 GQA,为什么没变快”。像预填充那样一次处理很多位置时,计算是瓶颈,而 g 对此几乎没有影响。收益来自带着长上下文逐个取出令牌的区间,以及提高并发请求数这件事。
第三,并发处理数的上限由缓存决定。把剩余内存除以一份缓存的大小,得到的值大致就是上限。把 key·value 头减少到四个,这一份就变小,所以上限提高。这时增加的是吞吐量,而不是单个请求的速度。
第四,在张量并行中,key·value 头数会成为绊脚石。把头分配到多个设备上时,如果 key·value 头比设备数少,就没有足够的东西可分,同样的 key·value 会被多个设备同时持有。为了缩小缓存而降低了 g,但每个设备都有一份副本的话,就达不到预期的缩减。
第五,改变头数,就是另一个模型。把取平均得到的权重原样部署到服务上,输出会不同。原论文在迁移之后也要做追加训练。并不是“只要改一下配置就行”。
实际工作中真正重要的事
- 减少的是内存。乘法次数几乎不变。要知道收益来自哪里,才能决定用在何处。
- h 必须能被 g 整除。如果各组大小不均匀,实现就无法成立。
- 两端在同一条轴上。g=h 就是 MHA,g=1 就是 MQA。如果把代码写成一份就能运行这三种,比较起来就很容易。
- 迁移与训练是不同的事。平均只是初始值。
下一项实验要做什么
把 /root/work/tf-gqa/gqa.py 一步一步做大。从把 query 头按连续的块分组开始,把组内的 key·value 头用平均折叠起来,再按 query 头的数量重复展开,然后经过单头的注意力,用同样的权重运行三种方式。把 g 设为 h,就必须与 MHA 一丝不差,把 g 减小,输出就必须不同。用数字测量这个差别。
接着是本实验的要点。分别数出 KV 缓存的元素个数和乘法次数,并排放在一起。缓存随 g 成正比缩小,而注意力主体的乘法一次也没有减少——这一点要靠自己数出来的数字来确认,而不是别人的说明。最后把这两张表和输出差别留成报告。
不测量时间。这个 Pod 没有 GPU,CPU 也与其他任务共用,所以“变快了”不是在这里能测的东西。系统 Python 里也没有 numpy(只在 /opt/onnx-lab 里有)。只用标准库就足够了,判定全部是元素个数、乘法次数这类整数计数,以及设置了容许误差的数值对照。评分器会真正导入你的模块,每次用不同的头数和组数调用函数,并与另行计算的值对照。