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

Transformer — 手算一遍注意力

亲手数出上下文长度的代价

在 TT Lab 中继续学习

目标

做一个数乘法的计数器,当上下文长度 n 增加时,亲手数出什么按 n 的平方增长,什么按 n 成正比增长。做出把两部分并排放置的表,并用自己数出的数字找出 n 的平方部分追上线性部分的交叉点。还要用整数数出:因果掩码留下的分数只有 n(n+1)/2 个,把分数矩阵整个放在手里时元素有 n 的平方个,追加一个令牌时新产生的分数只有一行。

为什么重要

“对上下文长度是平方”这句话只对了一半。一层里混着性质不同的两种计算。注意力的分数和混合在每个位置都要扫过整个上下文,所以按平方增长,而 query、key、value 投影、输出投影和前馈是对每一个位置都同样运行一次的工作,所以只与位置数成正比。哪边占主导,由长度和模型宽度共同决定。所以用一个倍率去估算增加上下文上限时的成本,必然会错。 本实验不测量时间。同样的代码,也会因机器和负载而不同,在容器里抖动得更厉害,测了也无法比较。取而代之的是数乘法次数。次数在哪里运行都是同样的整数,可以作为依据拿给别人看。这里出现的数字全都是你的计数器数出来的值,不测实际模型的秒数或 GB,所以也不使用。 模型的形状把 Attention Is All You Need 的 base 配置作为假设使用——D_MODEL = 512,D_FF = 2048,N_HEADS = 8。这是本实验定的假设,不是你在用的模型的值。 评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的大小检验函数,连计数器实际增加了多少都会对照。大小每次运行都会变,所以无法把值背下来填进去。

步骤

  1. 在 /root/work/tf-cost/cost.py 中创建常量 D_MODEL = 512、D_FF = 2048、N_HEADS = 8、计数器 MulCount,以及使用它的 dot(a, b, ctr)。
  2. 增加 attn_scores(Q, K, ctr)、attn_mix(A, V, ctr)、quad_mults(n, d_model),数出按 n 的平方增长的部分。
  3. 增加 matvec(M, x, ctr)、per_position_mults(d_model, d_ff, ctr)、linear_mults(n, d_model, d_ff),数出与位置数成正比的部分。
  4. 创建 cost_table(ns, d_model, d_ff),对每个长度返回 (n, n제곱 몫, 선형 몫, 합)(占位符依次为 n、n 的平方部分、线性部分与总和)。
  5. 创建 crossover_n(d_model, d_ff),找出 n 的平方部分第一次追上线性部分的长度。
  6. 创建 causal_pairs(n) 和 wasted_pairs(n),数出因果掩码留下的分数和丢弃的格子。
  7. 创建 score_bytes(n, n_heads, itemsize)、max_context_for_bytes(budget_bytes, n_heads, itemsize)、append_pairs(n)、generate_pairs(n, g)。
  8. 把实际运行上面这些函数的结果,记录到 /root/work/tf-cost/cost_report.json 和 /root/work/tf-cost/cost_report.md 中。

参考

不测量什么

不测量时间。也不使用实际模型的 GB 或秒数。把没有测过的数字写进记录,那份记录就没有依据。

做出数乘法的尺子

在 /root/work/tf-cost/cost.py 中创建常量 D_MODEL = 512、D_FF = 2048、N_HEADS = 8、计数器 MulCount(带有 mults、calls,用 add(k) 增加),以及 dot(a, b, ctr)。dot 返回点积的值,并把计数器恰好增加向量长度那么多。

不要去测量时间——同样的代码也会因机器和负载而不同,没法比较。长度为 d 的点积,乘法恰好是 d 次,所以一行 ctr.add(len(a)) 就够了。如果把一次调用数作 1,后面所有的数字都会崩。calls 是数分了几次记入账簿的值,在第 3 步会用到。

按 n 的平方增长的部分

增加 attn_scores(Q, K, ctr)、attn_mix(A, V, ctr)、quad_mults(n, d_model)。attn_scores 返回 n x n 的分数矩阵,attn_mix 返回用 A 混合 V 得到的 n x d,两者都让计数器增加 n * n * d。quad_mults 是两者之和 2 * n * n * d_model。

attn_scores 只要让每个 query 与全部 key 做 dot 就行。attn_mix 是容易搞混的地方——输出的 n x d 虽小,但每个输出位置都要扫过整个上下文的 n 个,所以乘法与分数计算一样,是 n * n * d 次。把 V 的竖列取出来交给 dot,计数器自然就对得上。quad_mults 里不含头数。

只与位置数成正比的部分

增加 matvec(M, x, ctr)、per_position_mults(d_model, d_ff, ctr)、linear_mults(n, d_model, d_ff)。per_position_mults 把六个线性映射(四个投影、两个前馈)各自真正运行一遍来数乘法,并返回增加的量。linear_mults 是那个值乘以位置数得到的整数。

矩阵里的值无所谓——目的是数数,所以做出形状对得上的矩阵运行就行。六个是四个 d_model x d_model、一个 d_ff x d_model、一个 d_model x d_ff。传给前馈第 2 层的向量,是第 1 层返回的长度为 d_ff 的向量。评分器还会检查计数器是否至少被调用了六次以上,把一个式子一次加完就结束的话,会被刷下来。linear_mults 不接收计数器。

把两部分并排放置

创建 cost_table(ns, d_model, d_ff)。对 ns 中的每个长度返回 (n, quad, linear, quad + linear) 四格,顺序与 ns 相同。

直接使用前面做的 quad_mults 和 linear_mults,五行就够了。请把长度每次翻倍地来读——前面的格子每次变四倍,后面的格子每次变两倍。最后一格必须是前两格之和。只写一侧,后面找交叉点时就会对不上。

找出交叉点

创建 crossover_n(d_model, d_ff)。是使 quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff) 第一次为真的 n。包含等号。

可以从 1 开始逐步增加去找,也可以用手解。把两边除以 n 和 2 * d_model,条件会变得非常短。漏掉等号只写不等号,答案会恰好推后一格。请改变模型宽度来调用——宽度越宽交叉点越往后推,这一点用数字就能看到。

数出掩码丢掉的一半

创建 causal_pairs(n) 和 wasted_pairs(n)。causal_pairs 是因果掩码中实际用到的分数个数,wasted_pairs 是 n * n 减去那一部分的值。

第 i 个 query 只看到自己为止。行的长度按 1、2、3 增加,到 n 结束,所以数出来就是 n * (n + 1) / 2。不能去掉对角线——看到自己才是对的。causal_pairs(0) 是 0,causal_pairs(1) 是 1。请打印几个点,看看被丢弃的比例随长度变长趋近于什么。

内存与一行一行增加的分数

创建 score_bytes(n, n_heads, itemsize)、max_context_for_bytes(budget_bytes, n_heads, itemsize)、append_pairs(n)、generate_pairs(n, g)。前两个是把分数矩阵整个放在手里时的字节数,以及预算内能放下的最长上下文,后两个是追加一个令牌时新产生的分数个数,和生成 g 个期间的总和。

score_bytes 容易漏掉头数。max_context_for_bytes 把实数平方根四舍五入会多出一格,所以请用 math.isqrt——把预算除以头数和字节数之后,取整数平方根即可。append_pairs(n) 因为新的 query 连自己在内看 n + 1 个,所以是 n + 1。generate_pairs 是从 n+1 到 n+g 的和,请务必确认 generate_pairs(0, N) 与 causal_pairs(N) 是否相同。

把数出的数字留成记录

实际运行上面这些函数,在 /root/work/tf-cost/cost_report.json 中写入 d_model、d_ff、n_heads、itemsize、table、quad_ratio、linear_ratio、crossover_n、quad_at_crossover、linear_at_crossover、causal_n、causal_pairs、wasted_pairs、wasted_fraction、score_bytes_at_causal_n、max_context_1gib、append_pairs_at_causal_n、generate_pairs_at_causal_n、identity_ok,在 /root/work/tf-cost/cost_report.md 中用 ## 무엇을 세었나(韩文,意为“数了什么”)、## 두 배로 늘리면 무엇이 네 배가 되나(韩文,意为“翻倍时什么会变成四倍”)、## 교차점은 어디인가(韩文,意为“交叉点在哪里”)、## 인과 마스크가 버리는 절반(韩文,意为“因果掩码丢掉的一半”)、## 메모리와 한 토큰씩 늘어나는 점수(韩文,意为“内存与每个令牌增加的分数”)五节来记录。

数字不要手写,要用运行你的代码得到的值来填。表的长度是 [128, 256, 512, 1024, 2048, 4096],quad_ratio、linear_ratio 是最后两行之间的倍率。causal_n 是 2048,数据类型是 2 字节,预算是 1073741824,生成令牌数是 256。identity_ok 是 generate_pairs(0, causal_n) == causal_pairs(causal_n) 是否为真。记录里不要写没有测过的数字——实际模型的秒数或 GB,这里从未测过。