亲手数出上下文长度的代价
目标
做一个数乘法的计数器,当上下文长度 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。这是本实验定的假设,不是你在用的模型的值。
评分器不会相信你写下的说明。它会真正导入你的模块,每次用不同的大小检验函数,连计数器实际增加了多少都会对照。大小每次运行都会变,所以无法把值背下来填进去。
步骤
- 在 /root/work/tf-cost/cost.py 中创建常量
D_MODEL = 512、D_FF = 2048、N_HEADS = 8、计数器MulCount,以及使用它的dot(a, b, ctr)。 - 增加
attn_scores(Q, K, ctr)、attn_mix(A, V, ctr)、quad_mults(n, d_model),数出按 n 的平方增长的部分。 - 增加
matvec(M, x, ctr)、per_position_mults(d_model, d_ff, ctr)、linear_mults(n, d_model, d_ff),数出与位置数成正比的部分。 - 创建
cost_table(ns, d_model, d_ff),对每个长度返回(n, n제곱 몫, 선형 몫, 합)(占位符依次为 n、n 的平方部分、线性部分与总和)。 - 创建
crossover_n(d_model, d_ff),找出 n 的平方部分第一次追上线性部分的长度。 - 创建
causal_pairs(n)和wasted_pairs(n),数出因果掩码留下的分数和丢弃的格子。 - 创建
score_bytes(n, n_heads, itemsize)、max_context_for_bytes(budget_bytes, n_heads, itemsize)、append_pairs(n)、generate_pairs(n, g)。 - 把实际运行上面这些函数的结果,记录到 /root/work/tf-cost/cost_report.json 和 /root/work/tf-cost/cost_report.md 中。
参考
- 执行契约:评分器会把
/root/work/tf-cost/cost.py当作 Python 模块导入,直接使用D_MODEL、D_FF、N_HEADS、MulCount、dot、attn_scores、attn_mix、quad_mults、matvec、per_position_mults、linear_mults、cost_table、crossover_n、causal_pairs、wasted_pairs、score_bytes、max_context_for_bytes、append_pairs、generate_pairs。它不会作为脚本运行,所以可以没有if __name__ == "__main__"。 MulCount带有mults(到目前为止数出的乘法次数)和calls(记入账簿的次数)两个值,用add(k)加上 k。一开始两者都是 0。dot(a, b, ctr)返回点积的值,并把计数器恰好增加长度那么多。是把一次乘法数作 1,而不是把一次调用数作 1。attn_scores(Q, K, ctr)返回 Q 的行数 x K 的行数大小的矩阵。Q 和 K 各自是 n x d 时,计数器会增加n * n * d。attn_mix(A, V, ctr)在 A 是 n x n、V 是 n x d 时返回 n x d,计数器又增加n * n * d。不能因为输出小,计算就小。quad_mults(n, d_model)是把两部分加起来的2 * n * n * d_model。不包含头数——因为每个头的宽度减为d_model / h,而这样的头有 h 个,所以总和相同。per_position_mults(d_model, d_ff, ctr)要把六个线性映射各自真正运行一遍来数乘法。query、key、value 三个投影(d_model x d_model),一个输出投影(d_model x d_model),前馈第 1 层(d_ff x d_model)和第 2 层(d_model x d_ff)。返回的值是这个函数增加的量,评分器还会检查计数器是否至少被调用了六次以上。把一个式子一次加完就结束的话,会被刷下来。linear_mults(n, d_model, d_ff)是一个位置的值乘以位置数得到的整数。不接收计数器。cost_table(ns, d_model, d_ff)返回的每一行有(n, quad, linear, quad + linear)四格,顺序与ns相同。crossover_n(d_model, d_ff)是使quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff)第一次为真的 n。包含等号。可以从 1 开始逐步增加去找,也可以解式子。causal_pairs(n)由“第 i 行用i + 1个”这一事实得出。causal_pairs(0)是 0,causal_pairs(1)是 1。wasted_pairs(n)是n * n减去实际用到的部分。score_bytes(n, n_heads, itemsize)是n * n * n_heads * itemsize。别漏掉头数。max_context_for_bytes(budget_bytes, n_heads, itemsize)是满足score_bytes(n, ...) <= budget_bytes的最大 n。把实数平方根四舍五入会多出一格,所以请用math.isqrt。append_pairs(n)是n + 1,generate_pairs(n, g)是行长度从n+1到n+g的和。generate_pairs(0, N)必须与causal_pairs(N)相同。- 第 8 步的报告使用
D_MODEL、D_FF、N_HEADS和ITEMSIZE = 2(假设是两字节的数据类型)、预算1073741824(1 GiB)、表的长度[128, 256, 512, 1024, 2048, 4096]、因果与内存计算的基准长度2048、生成令牌数256。 quad_ratio、linear_ratio是表的最后两行(2048 和 4096)之间的倍率。因为是除法,所以是实数,评分器用abs(a - b) <= atol + rtol * abs(b)来比较。- 这个 Pod 没有互联网。
pip install无法使用,系统 Python 里没有 numpy、torch、transformers。numpy 只在/opt/onnx-lab/bin/python里。只用标准库就足够了。 - 官方文档:Attention Is All You Need · PyTorch — scaled_dot_product_attention · Hugging Face — Cache strategies
- 常见错误:把计数器按调用次数数、以为
attn_mix只数输出那么大的量、在线性部分漏掉输出投影、找交叉点时不带等号而推后一格、在causal_pairs里去掉对角线、score_bytes里没有乘头数。
不测量什么
不测量时间。也不使用实际模型的 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,这里从未测过。