上下文翻倍,什么会变成四倍
一句话总结
Transformer 一层的计算分成两部分。注意力的分数和混合随上下文长度的平方增长,其余的线性变换随长度成正比增长。在短上下文中后者占主导,超过某个长度之后前者占主导——弄清楚这个位置,就是本文的全部。
为什么需要它
增加上下文上限的请求,哪个团队都会遇到。从 4 千增加到 8 千,从 8 千增加到 3 万 2 千。被问到增加之后会贵多少,回答“大概两倍”,看到账单又吃惊,这种事反复发生。
反过来的情况也有。上下文翻了一倍,成本却只增加了两倍,于是想“不是说平方吗,也没有嘛”就放过去了。两者出自同一个误解。只看了其中一边。
一层里混着性质不同的两种计算。注意力在每个位置都要扫过整个上下文,所以位置数一增加,扫描的工作也跟着增加,成为平方。而生成 query、key、value 的投影、输出投影、前馈(feed-forward),是对每一个位置都同样运行一次的工作,所以只与位置数成正比。
所以“对上下文长度是平方”这句话只对了一半。准确地说是按平方增长的部分与按正比增长的部分之和,哪边占主导,由长度和模型宽度共同决定。
该数什么
这里不能测量时间。同样的代码,也会因机器和负载而不同,在容器里抖动得更厉害。所以数乘法次数。次数在哪里运行都是同样的整数。
数法很简单。两个长度为 d 的向量的点积,乘法恰好是 d 次。其余只要数这个点积出现了几次。
def dot(a, b, ctr):
total = 0.0
for x, y in zip(a, b):
total += x * y
ctr.add(len(a)) # 곱셈 d 번을 장부에 적는다
return total
用这把尺子去量一量注意力的定义,结果是这样。
- 分数矩阵 QK:有 n 行、n 列,每个格子是长度为 d 的点积 →
n * n * d - 混合 AV:输出是 n x d,虽然小,但每个位置都要扫过整个上下文的 n 个 → 又是
n * n * d - 即使把头拆成 h 个,总和也一样。因为每个头的宽度减为
d / h,而这样的头有 h 个。
所以按平方增长的部分是 2 * n * n * d。另一边,只要把一个位置所付出的值数一次,再乘以位置数。
- query、key、value 三个投影加一个输出投影 → 每个位置
4 * d * d - 前馈两层 → 每个位置
2 * d * d_ff
合起来每个位置是 4 * d * d + 2 * d * d_ff,位置有 n 个。
交叉点
把两个式子并排放,交叉点用手就能解出来。只要找出使 2 * n * n * d 大于等于 n * (4 * d * d + 2 * d * d_ff) 的最小 n。两边除以 2 * n * d,条件就化简为 n >= 2 * d + d_ff。
数字带来的感觉很重要。模型越宽,交叉点就越往后推。因为一个位置所付出的线性成本随宽度的平方增大,而注意力一侧只随宽度成正比增大。在大模型里“不是说平方吗,怎么没感觉”的体感就来自这里——还在交叉点之前。
而一过交叉点,情况就不同了。把长度翻一倍,总和开始趋近四倍。
因果掩码丢掉的一半
解码器的注意力有因果掩码。第 i 个位置只看到自己,所以行的长度按 1、2、3 增加,到 n 结束。实际用到的分数是 n * (n + 1) / 2 个。
矩阵仍然是 n 的平方格。也就是说,近一半是算出来再丢掉的。准确地说,被丢掉的是 (n - 1) / (2 * n),长度越长越趋近一半。
天真地全部算出来再盖上掩码的实现,就是这样做的。要减少计算,就不是事后再盖掩码,而是根本不去计算,于是出现了按块只运算三角形部分的内核(kernel)。PyTorch 的 scaled_dot_product_attention 单独接收 is_causal,也是同样的原因。把掩码作为张量传入再相乘,与告诉它“这是因果的”,内部运行的东西不同。
内存也是同样的形状
把分数矩阵整个放在手里,每个头的值的个数是 n 的平方。合起来是 n * n * h 个,再乘以一个数据类型格子的字节数,就得到字节数。
这为什么痛,反过来看就知道了。定下预算,求能放进去的最大长度,会发现预算要增加四倍,长度才能翻倍。意思是把显卡增加一倍,上下文只能增加 1.41 倍。
所以不整个生成分数矩阵的实现变得重要。切成块,一块一块地处理再丢掉,做同样的计算,手里拿着的值的数量却减少了。计算量不变,只是内存减少——如果没有这两者各行其是的感觉,就无法理解“为什么计算一样,上下文却能更长”。
放入的长度与生成的长度,增长方式不同
一次性放入提示词(预填充,prefill)和一个一个地生成令牌(token)(解码,decode),增长的形状不同。
预填充要一次处理 n 个位置,所以分数以 n 的平方的规模产生。而在已经积累了 n 个之后再多加一个,新的 query 只是看 n + 1 个 key。只有一行。所以生成 g 个的过程中产生的分数总和是 g * n + g * (g + 1) / 2。
这里会引出一个有趣的恒等式。把 n 设为 0,这个值与 causal_pairs(g) 恰好相同。无论是一次性放进去再用掩码擦掉,还是一个一个地追加,实际需要的分数个数是一样的。不同的只是一次做还是分开做。用什么来弥补这个差别,是 KV 缓存文档所讨论的主题,属于下一个模块。
在现场相遇的样子
第一,把上下文上限翻倍,延迟却没有翻倍,于是放心了。只是因为还在交叉点之前,线性部分占主导。再增加长度,斜率会突然改变。
第二,同样的长度,换了模型,成本曲线的形状就不同。宽度不同,交叉点在不同的位置。把在一个模型上测得的倍率原样用在另一个模型上,会出错。
第三,长上下文下内存先爆。计算还撑得住,但放分数矩阵的地方没有了。计算量和内存即使同样是 n 的平方,先撞上的墙通常是内存。
第四,把掩码做成张量再相乘的实现慢了一倍。因为把全部算出来又丢掉了一半。同样的公式,什么时候用掩码,决定了实际的工作量。
第五,预填充很慢,生成令牌却很快。或者相反。两者的增长形状不同,所以用一个倍率捆在一起估算,必然有一边是错的。
实际工作中真正重要的事
- 不要只背“平方”这个词,要拆成两部分来看。哪边占主导,由长度和宽度共同决定。
- 数次数,不要数时间。次数换了机器也一样,向别人解释估算时可以作为依据。
- 用自己模型的数字求出交叉点。在那个点的前后,容量规划是不同的。
- 把计算量和内存分开看。即使同样是 n 的平方,减少的办法也不同。
- 把预填充和解码分开测量。一个倍率必然会让一边出错。
下一项实验要做什么
一步步扩展 /root/work/tf-cost/cost.py。只使用标准库——这个 Pod 的系统 Python 里没有 numpy(只在 /opt/onnx-lab/bin/python 里才有),也没有 torch 和 transformers。取而代之的是亲手做一个数乘法的计数器,只使用由它得出的整数。
从计数器和点积开始,真正运行分数矩阵和混合,数出 n 的平方来加以确认。接着真正运行一个位置所付出的六个线性变换,确认这一部分与长度无关,并做出把两部分并排放置的表。
在那张表里找出交叉点。还要改变模型宽度,用自己的数字看到交叉点移到了哪里。接着数出因果掩码留下的分数,确认有一半被丢掉,并求出整个放进分数矩阵时的字节数以及预算内能放下的最大长度。
最后一步是本实验的要点。数出追加一个令牌时新产生的分数只有一行,并用整数确认:从长度 0 生成 g 个时的总和,与因果掩码留下的个数恰好相同这个恒等式。评分器会真正导入你的模块,每次用不同的大小检验函数,连计数器实际增加了多少都会对照。值无法背下来填进去。