TT Lab
开始
学习 学习路径 课程

Transformer — 手算一遍注意力

只把键值头减下来

在 TT Lab 中继续学习

目标

只用标准库做出保持 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 合并起来,答案就会不同这个事实的值。 评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的头数和组数调用函数,并与评分器另行计算的值对照。

步骤

  1. 在 /root/work/tf-gqa/gqa.py 中创建 make_heads(count, rows, dim, seed) 以及 group_of(head, h, g)、group_members(h, g)。样本在参数相同时必须永远得到同样的值,组要分成连续的块。
  2. 增加 fold_kv(heads, g),把组内的 key·value 头逐位置取平均,折叠成 g 个。
  3. 增加 expand_kv(folded, h),把折叠起来的头就地重复,展开成 h 个。
  4. 增加 attend(q, k, v)。把分数除以 sqrt(헤드 차원)(占位符为头维度),softmax 要先减去最大值再取指数。
  5. 增加 heads_out(q_heads, k_heads, v_heads, g) 和 max_gap(left, right),用同样的权重运行三种方式并测量输出差别。
  6. 增加 kv_cache_elems(layers, kv_heads, seq_len, head_dim) 和 kv_table(layers, h, groups, seq_len, head_dim),数出缓存的元素个数。
  7. 增加 mults(n, d_model, h, g, head_dim),分项数出一层通过一次的乘法次数。
  8. 定好样本和模型规模,做出两张表,并把结果记录到 /root/work/tf-gqa/gqa_report.json 和 /root/work/tf-gqa/gqa_report.md 中。

参考

把 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 中把缓存的倍数和乘法的倍数并排写下来——一边最多缩小到八分之一,另一边几乎不变,这就是本实验的结论。