只把键值头减下来
目标
只用标准库做出保持 query 头数不变、只减少 key·value 头数的结构。把 query 头按连续的块分组,把组内的 key·value 头用平均折叠起来,再按 query 头的数量重复展开,然后运行注意力。仅仅改变组数 g,就用同样的权重运行 MHA(g=h)、GQA(1
为什么重要
如今的模型配置里,query 头数和 key·value 头数是分开写的。两个值相同就是 MHA,key·value 一侧是 1 就是 MQA,介于两者之间就是 GQA。比起背这三个名字,更重要的是知道什么在减少。
减少的是 KV 缓存的内存。缓存的元素个数是 2 x 층 x 키·값 헤드 수 x 길이 x 헤드 차원(韩文,意为“2 x 层数 x key·value 头数 x 长度 x 头维度”),这个式子里没有 query 头数。而注意力主体的乘法次数是 h x n x n x 헤드차원(韩文,意为“h x n x n x 头维度”),所以组数 g 根本不会进来。因为要把折叠起来的 key·value 按 query 头的数量重复展开,然后照常运行。亲手数一数这两个式子,就能回答“换成了 GQA,为什么预填充还是老样子”这个问题。
本实验不测量时间。Pod 里没有 GPU,CPU 也与其他任务共用,所以在这里测出的速度什么也说明不了。判定全部是元素个数、乘法次数这类整数计数,以及设置了容许误差的数值对照。
把组内的 key·value 头用平均来合并,是本实验的假设。它与 GQA 原论文迁移已训练好的多头检查点(checkpoint)时用的方式相同,但原论文在平均之后还要做追加训练。这里得到的输出差别,并不表示质量会变差这么多,而是显示在同样的权重下只把 key·value 合并起来,答案就会不同这个事实的值。
评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的头数和组数调用函数,并与评分器另行计算的值对照。
步骤
- 在 /root/work/tf-gqa/gqa.py 中创建
make_heads(count, rows, dim, seed)以及group_of(head, h, g)、group_members(h, g)。样本在参数相同时必须永远得到同样的值,组要分成连续的块。 - 增加
fold_kv(heads, g),把组内的 key·value 头逐位置取平均,折叠成 g 个。 - 增加
expand_kv(folded, h),把折叠起来的头就地重复,展开成 h 个。 - 增加
attend(q, k, v)。把分数除以sqrt(헤드 차원)(占位符为头维度),softmax 要先减去最大值再取指数。 - 增加
heads_out(q_heads, k_heads, v_heads, g)和max_gap(left, right),用同样的权重运行三种方式并测量输出差别。 - 增加
kv_cache_elems(layers, kv_heads, seq_len, head_dim)和kv_table(layers, h, groups, seq_len, head_dim),数出缓存的元素个数。 - 增加
mults(n, d_model, h, g, head_dim),分项数出一层通过一次的乘法次数。 - 定好样本和模型规模,做出两张表,并把结果记录到 /root/work/tf-gqa/gqa_report.json 和 /root/work/tf-gqa/gqa_report.md 中。
参考
- 执行契约:评分器会把
/root/work/tf-gqa/gqa.py当作 Python 模块导入,直接使用make_heads、group_of、group_members、fold_kv、expand_kv、attend、heads_out、max_gap、kv_cache_elems、kv_table、mults。它不会作为脚本运行,所以可以没有if __name__ == "__main__"。 make_heads(count, rows, dim, seed)返回 count 个头,每个头有 rows 行,一行是 dim 个实数。参数相同时必须永远得到同样的值(禁止random模块),seed 不同时必须得到不同的值。值的绝对值取在 4 以下,不能所有值都相同。group_of(head, h, g)按连续的块来分配。h=8、g=2 时,0、1、2、3 是第 0 组,4、5、6、7 是第 1 组。如果h % g != 0,或者 h、g 小于 1,或者 head 超出范围,就抛出ValueError。group_members(h, g)是长度为 g 的列表,每一格是属于该组的 query 头编号列表。fold_kv(heads, g)把组内的头逐位置取平均,折叠成 g 个。g 等于头数时,值原样输出(每组只有一个头,平均就是它自己)。如果不能均匀分开,就是ValueError。expand_kv(folded, h)把每个头就地重复。把[A, B]展开成 4 个是[A, A, B, B],而不是[A, B, A, B]。如果不能均匀展开,就是ValueError。attend(q, k, v)返回与 q 形状相同的表。不使用掩码——这里要看的只是 key·value 头数的效果,所以去掉因果掩码来比较。heads_out(...)用fold_kv折叠、用expand_kv展开之后,对每个 query 头各调用一次attend。用g = len(q_heads)调用时,折叠和展开不会改变值,所以必须得到与普通多头相同的结果。max_gap(left, right)返回两个输出之间最大的绝对值之差,一个数。kv_cache_elems是2 * layers * kv_heads * seq_len * head_dim。不包含 query 头数。kv_table是[(무리 수, 원소 수), ...](占位符依次为组数与元素个数)。mults(n, d_model, h, g, head_dim)的键有proj_q、proj_k、proj_v、scores、weighted、proj_out、total七个,值全是整数。total是其余六个之和。只数乘法,不数加法和指数。- 第 8 步报告的值这样来定。输出比较用
sample_h = 8、sample_n = 6、sample_head_dim = 4,并把make_heads(8, 6, 4, 101)、make_heads(8, 6, 4, 202)、make_heads(8, 6, 4, 303)分别用作 Q、K、V。组数是groups = [8, 4, 2, 1]。 - 模型规模取
layers = 32、seq_len = 4096、head_dim = 128、h = 8、d_model = 1024、n_tokens = 4096。这不是从实际模型上测来的值,而是为了做表而设的示例规模。 - 要放进
gqa_report.json的键:sample_h、sample_n、sample_head_dim、groups、diff_table、layers、seq_len、head_dim、h、d_model、n_tokens、kv_table、kv_ratio、mult_table、mult_ratio、core_mults、core_same_for_all_g。diff_table是[[무리 수, MHA 와의 최대 차이], ...](占位符依次为组数与同 MHA 的最大差别),kv_ratio、mult_ratio是把 MHA 设为 1 的倍数,即MHA 값 / 그 무리의 값(韩文,意为“MHA 的值 / 该组的值”),core_mults是scores + weighted。 gqa_report.md用## 무엇을 쟀나(韩文,意为“测量了什么”)、## 무리를 줄이면 출력이 얼마나 달라지나(韩文,意为“减少组数后输出会变化多少”)、## 메모리는 줄고 곱셈은 안 준다(韩文,意为“内存减少而乘法不减少”)、## 어디에 쓸 것인가(韩文,意为“用在哪里”)四节来写。- 这个 Pod 没有互联网,也没有 GPU。
pip install无法使用,numpy 只在/opt/onnx-lab/bin/python里有,所以在系统 Python 中import numpy不可用。只用math就足够了。 - 官方文档:MQA 原论文 · GQA 原论文 · Attention Is All You Need · PyTorch — MultiheadAttention
- 常见错误:把组交替分配;折叠时只留下第一个头;展开时把整个列表接在一起;分数不做除法;softmax 里没有减去最大值;把 query 头数放进缓存的式子;把乘法式子里的
h换成g,以为计算量也会减少。
把 query 头分成组
在 /root/work/tf-gqa/gqa.py 中创建 make_heads(count, rows, dim, seed) 以及 group_of(head, h, g)、group_members(h, g)。样本在参数相同时必须永远得到同样的值(禁止 random 模块,绝对值在 4 以下),组要分成连续的块。h=8、g=2 时,0、1、2、3 是第 0 组。如果不能均匀分开,就抛出 ValueError。
样本用一个小的线性同余式就够了——用整数保存状态,重复 state = (a * state + c) % m,同时取出 state / m - 0.5,同一个 seed 就会得到同样的数列。组编号是 head // (h // g)。head % g 是交替分配,到后面展开时会错位。group_members 如果通过调用 h 次 group_of 来填充,两处的规则就不会产生分歧。
把组内的 key·value 折叠成一个
增加 fold_kv(heads, g)。把组内的 key·value 头逐位置取平均,折叠成 g 个。如果头数不能被 g 整除,就是 ValueError。当 g 等于头数时,值必须原样输出。
组的划分与第 1 步的规则相同,也就是连续的块。用 heads[start:start + size] 切出一个组,逐位置相加后除以组的大小就行。如果只留第一个头,或者无视分组而全部取平均,只有在 g=1 时碰巧看起来相同。平均是本实验的假设——它与 GQA 原论文迁移检查点时用的方式相同,但原论文在那之后还要做追加训练。
再展开成 query 头的数量
增加 expand_kv(folded, h)。把折叠起来的头就地重复,展开成 h 个。把 [A, B] 展开成 4 个是 [A, A, B, B]。如果不能均匀展开,就是 ValueError。
外层循环是折叠后的头,内层循环是 h // len(folded) 次。如果把整个列表乘起来再接在一起(folded * size),就会变成 [A, B, A, B],使 query 头 i 看到别的组的 key·value。请确认它与第 1 步的 group_of 配得上——query 头 i 所看的必须是 folded[group_of(i, h, g)]。行用新列表复制后再返回,以后就不会出现改动一处而多个头一起变化的情况。
单头的注意力
增加 attend(q, k, v)。分数除以 sqrt(헤드 차원)(占位符为头维度),softmax 要先减去最大值再取指数。返回的值是与 q 形状相同的表。不使用掩码。
对每个 query 行,与所有 key 行做点积,相除,经过 softmax,再对 value 行做加权平均。如果不用 sqrt 去除,维度越大,softmax 越会偏向一个位置。如果不减去最大值,在分数很大的样本上 math.exp 会溢出,发生 OverflowError——评分器会故意放进很大的值试试。
用同样的权重运行三种方式
增加 heads_out(q_heads, k_heads, v_heads, g) 和 max_gap(left, right)。heads_out 把 key·value 折叠成 g 个再展开成 h 个之后,对每个 query 头各调用一次 attend。g = len(q_heads) 时必须与普通多头一丝不差,把 g 减小,输出就必须不同。max_gap 返回两个输出之间最大的绝对值之差。
三行就结束——把 key 折叠再展开,把 value 折叠再展开,对每个 query 头调用 attend。不碰 query 头,就是这一步的全部。max_gap 把头、行、列全部扫一遍,只留下最大的那一个差。一边改变 g 一边调用,可以看到 g=h 时差恰好是 0,而 g 越小,差越大。
数出缓存的元素个数
增加 kv_cache_elems(layers, kv_heads, seq_len, head_dim) 和 kv_table(layers, h, groups, seq_len, head_dim)。元素个数是 2 * layers * kv_heads * seq_len * head_dim,不包含 query 头数。kv_table 返回 [(무리 수, 원소 수), ...](占位符依次为组数与元素个数),如果有不能均匀分开的组数,就是 ValueError。
前面的 2 是 K 和 V 两份。有些地方会让人想乘上 query 头数,但 query 不会留在缓存里,所以式子里没有。在 kv_table 中,直接利用“该组数下的 key·value 头数就是 g”这一点就行。不按字节而按元素个数来数,是因为随数据类型不同,一个元素可能是 2 字节,也可能是 4 字节。
数乘法次数
增加 mults(n, d_model, h, g, head_dim)。键有 proj_q、proj_k、proj_v、scores、weighted、proj_out、total,值是整数。scores 和 weighted 是 h * n * n * head_dim,投影是 n * d_model * (헤드 수 * head_dim)(占位符为头数)这样的形式,total 是其余六个之和。如果不能均匀分开,就是 ValueError。
哪一项里有 g、哪一项里没有,就是这一步的全部。把 key·value 展开之后,注意力按 query 头的数量运行,所以 scores 和 weighted 里没有 g。与 g 有关的只有 K 投影和 V 投影两个。只数乘法,不数加法和 softmax 的指数——要看的不是总量,而是哪一项在减少。
记录什么减少、什么不减少
样本用 make_heads(8, 6, 4, 101)、make_heads(8, 6, 4, 202)、make_heads(8, 6, 4, 303) 分别作为 Q、K、V,组数取 [8, 4, 2, 1]。模型规模是 layers = 32、seq_len = 4096、head_dim = 128、h = 8、d_model = 1024、n_tokens = 4096。在 /root/work/tf-gqa/gqa_report.json 中写入 sample_h、sample_n、sample_head_dim、groups、diff_table、layers、seq_len、head_dim、h、d_model、n_tokens、kv_table、kv_ratio、mult_table、mult_ratio、core_mults、core_same_for_all_g,并在 /root/work/tf-gqa/gqa_report.md 中用 ## 무엇을 쟀나(韩文,意为“测量了什么”)、## 무리를 줄이면 출력이 얼마나 달라지나(韩文,意为“减少组数后输出会变化多少”)、## 메모리는 줄고 곱셈은 안 준다(韩文,意为“内存减少而乘法不减少”)、## 어디에 쓸 것인가(韩文,意为“用在哪里”)四节来写。
数字不要手写,要用实际运行你的函数得到的值来填。diff_table 对每个组数是 max_gap(그 무리의 출력, g=h 의 출력)(韩文,意为“max_gap(该组的输出, g=h 的输出)”),所以第一格是 0.0。kv_ratio、mult_ratio 是把 MHA 设为 1 的倍数,所以是 MHA 값 / 그 무리의 값(韩文,意为“MHA 的值 / 该组的值”)。core_mults 是 scores + weighted,组数变化时值也相同,所以 core_same_for_all_g 为真。请在 md 中把缓存的倍数和乘法的倍数并排写下来——一边最多缩小到八分之一,另一边几乎不变,这就是本实验的结论。