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

Transformer — 手算一遍注意力

上下文翻倍,什么会变成四倍

在 TT Lab 中继续学习

一句话总结

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

用这把尺子去量一量注意力的定义,结果是这样。

所以按平方增长的部分是 2 * n * n * d。另一边,只要把一个位置所付出的值数一次,再乘以位置数。

合起来每个位置是 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。

把一层的乘法次数拆成两部分画出的图。注意力的分数和混合是随上下文长度的平方增长的曲线,其余的线性变换是与长度成正比的直线。两条线在 n 等于 2d 加 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 的平方,先撞上的墙通常是内存。

第四,把掩码做成张量再相乘的实现慢了一倍。因为把全部算出来又丢掉了一半。同样的公式,什么时候用掩码,决定了实际的工作量。

第五,预填充很慢,生成令牌却很快。或者相反。两者的增长形状不同,所以用一个倍率捆在一起估算,必然有一边是错的。

实际工作中真正重要的事

下一项实验要做什么

一步步扩展 /root/work/tf-cost/cost.py。只使用标准库——这个 Pod 的系统 Python 里没有 numpy(只在 /opt/onnx-lab/bin/python 里才有),也没有 torch 和 transformers。取而代之的是亲手做一个数乘法的计数器,只使用由它得出的整数。

从计数器和点积开始,真正运行分数矩阵和混合,数出 n 的平方来加以确认。接着真正运行一个位置所付出的六个线性变换,确认这一部分与长度无关,并做出把两部分并排放置的表。

在那张表里找出交叉点。还要改变模型宽度,用自己的数字看到交叉点移到了哪里。接着数出因果掩码留下的分数,确认有一半被丢掉,并求出整个放进分数矩阵时的字节数以及预算内能放下的最大长度。

最后一步是本实验的要点。数出追加一个令牌时新产生的分数只有一行,并用整数确认:从长度 0 生成 g 个时的总和,与因果掩码留下的个数恰好相同这个恒等式。评分器会真正导入你的模块,每次用不同的大小检验函数,连计数器实际增加了多少都会对照。值无法背下来填进去。