自注意力 入门
把自注意力画到小得能看清:一个六词的句子,每个向量两维,所有查询、键、值都活在同一张平面上。分数矩阵、√d 除数、softmax、因果掩码和加权和 —— 页面上每一个数字都由旁边那张图算出来,包括那张演示掩码被写错时会发生什么的图。
一个向量,三个身份
一个 token 进来时只是一个向量。注意力做的第一件事,是把它拆成一个查询、一个键和一个值。
例句是 the cat ate and it ran,整篇文章跟着其中一个 token 走:代词 it。每个 token 画成二维向量,而不是真实注意力头用的 64 维,这样页面上每一个量都是平面上能用手指指出来的东西。论证本身跟宽度无关。
三个学到的矩阵把这一个向量变成三个:W_Q 给出这个 token 的查询,W_K 给出它的键,W_V 给出它的值。拖滑块,一个一个把它们装上去:
注意三支箭从同一个输入出发,却指向三个不同方向。查询是这个 token 在找什么,键是它对外宣告什么,值是别人真的选中它时它交出什么。名字借自信息检索,这个类比一路撑到本文结尾。
三个矩阵在这一层里是固定的,所以三支箭是输入的刚性函数:挪动 token,查询、键和值一起跟着挪。把灰色的输入向量拖到平面上任何地方:
看着输入往上推时查询甩得多厉害,而值几乎没转。这正是三个矩阵买到的东西:一个词要找什么、对外宣告什么、贡献什么,是它的三个不同函数,一个向量扛不动这三件事,硬扛就等于把它们变成同一件事。
其他每个 token 在同一平面上也有自己的键。逐个走过这六个 token,看每个查询把最长的影子投在哪个键上:
在 处,赢家是cat 的键,1.72 对 it 自己那个键的 1.14。没有人告诉模型 it 指的是 cat:梯度下降把 W_Q 和 W_K 塑造到代词的查询会偏向名词的键,因为这样损失才降得下来。
如果这两个映射是同一个映射,查询就是它自己的键,问别人和被别人问就成了同一回事。把 W_Q 混进 W_K,看这两个方向怎么并到一起:
没动过时,cat 给 ate 打 0.51,而 ate 给 cat 打 −0.32:动词要主语,远比主语要动词更迫切。到时两边都是 1.16,这段关系退化成了相似度。语言不是对称的,把两个映射分开,买到的正是方向。
分数就是“对得有多齐”
两个向量进去,一个数出来。这个数说的是查询沿着键伸出去多远。
算术上它是点积:对应槽位相乘,全部加起来。几何上它是一道影子。从查询的尖端向键所在的直线作垂线,分数就是这道影子的长度乘以键的长度。拖动查询,盯着那个数看:
注意符号在哪里翻转。查询和键朝同一侧时分数为正,恰成直角时为零,倒向另一侧后变负 —— 负分不是错误,而是这个 token 在明确地投反对票。
一个查询要同时对上全部六个键,于是这张平面变成一行六个数。让查询走过整句话,从这些柱子上把这一行读出来:
因为这六个分数是裸点积,它们两个方向都没有上界:在自己的键上冲到 3.02,而在 cat 上掉到 −0.45。到这一步为止,还没有任何东西说明谁算“大”。
裸点积的陷阱也正在这里:它是两个长度的乘积,所以一个键可以靠“长”取胜,而不是靠指向有用的方向。把 and 的键拉长,看它一路爬上来:
拉到 时,虚词 and 的分数就超过了 cat,而它一度都没转过。真实的 Transformer 靠投影之前的 LayerNorm 压住这件事,近来一些模型还直接对 q 和 k 做归一化(QK-norm),正是因为一个过长的键足以霸占矩阵的每一行。
每个查询都要碰上每个键
一个查询占一行,一个键占一列。这个方阵就是 Q · Kᵀ 的全部。
这里没有检索,也没有候选集。每个 token 的查询都要和每个 token 的键做点积,包括它自己那个,结果铺成一个边长等于句长的方阵。一次填一行,把它填满:
每个小方块的边长就是分数的大小,所以不用读数字也能看懂这张图。六个 token 要 ;一千个 token 的上下文要一百万次,每个头、每一层都要一遍。大家抱怨的那个平方复杂度就是这个 —— 不过在 GPT-2 small 里,四个投影的开销要一直大过它,直到 n 越过 2·d_model,也就是 1,536 个 token。
一行就是一个查询对上整句话 —— 也就是 §02 从柱子上读出来的那六个数,现在和其他人的叠在了一起。顺着行往下走:
最大的方块落在 cat 底下,1.22;而几乎是平的:一个限定词没有什么特别要找的东西,而一行几乎持平,正是注意力在说我没有偏好。
很容易把这个方阵当成相似度表,但它不是。沿对角线翻一下,看画面怎么变:
因为 W_Qᵀ W_K 不对称,cat → ate 的 0.51和 ate → cat 的 −0.32 是关于同一对词的两个不同的数。动词伸手去够它的主语,主语却给动词投了反对票。相似度矩阵分不出这两件事,而这份不对称,正是这里有两个映射而不是一个的全部理由。
从分数到一份混合
进去的是六个没有上下界的数,出来的是六个加起来等于一的正权重。
softmax 把每个分数取指数,再除以总和。取指数让每一项无论分数是正是负都变成正数,除以总和让这一行加起来等于一。青色的阶梯从左往右把这些权重累加起来,它必须正好落在上边缘:
注意这份混合有多“软”。在里,最大的权重是 cat 的 0.30,最小的是 the 的 0.11 —— 差三倍,而不是一次查表。注意力几乎从不只挑一个 token,它给每个都倒一点,给某一个多倒一点。
这六个不是六个独立的决定。它们共用同一个分母,所以抬高任何一个分数,必然压低其余每一个权重。沿着这一行拖动 cat 的分数,看另外五个怎么让位:
盯着右边那个读数:不管你把某一个分数怎么折腾,总和始终是 1.000。这就是整个机制的不变量,值得直接写成断言 —— 每一行都满足 all(w > 0) 且 sum(w) == 1,正是它让输出成为一个加权平均,而不是任意的线性组合。
这条断言的下界是严格的。分数想压多低都行,权重会一路趋近零,却永远到不了零。压下去,读那个指数:
在 时,权重是 1.2e-14 —— 柱子上看不见,算术里仍然是正的。这就是为什么掩码必须用 −∞ 而不是“一个很负的数”,也是为什么 softmax 可以安全地实现成 exp(s − max(s)):这个平移在比值里被约掉了,所以稳定版本就是同一个函数,而不是它的近似。
公式里那个 √d 是干什么的
这个除数不是凑出来的系数,它就是被除的那个量的离散度。
两个 d 维向量的点积,是 d 个乘积之和。如果各分量独立、均值为 0、方差为 1,那每个乘积方差是 1,和的方差就是 d —— 于是分数的离散度随 √d 增长。拖动宽度,从曲线上把它读出来:
宽度轴是对数轴 —— 每走一格宽度就是四倍 —— 而离散度每格仍然翻一番,所以这条曲线即便在这里也是往上弯的。位于 1 的那条琥珀色横线,是 softmax 正常工作所需要的离散度。在 处,离散度是 8.00 —— 宽了八倍。
宽八倍不等于噪声大八倍。softmax 是指数函数,所以把每个分数乘以八,等于把最大的那个相对于对手抬到了八次方。把宽度推上去,看这份混合怎么塌掉:
到 时,一个权重已经占掉 0.991,到 64 就全归它了。这比“随便选个数”更糟:softmax 经由某个权重回传的梯度正比于 w(1 − w),所以一行被钉在 1.000 上,几乎什么都传不回去,也就不再学习它本该选谁。
把每个分数都除以 √d,无论宽度多大,离散度都被放回到 1。把除数装上,再拉同一个滑块:
这份混合纹丝不动:从 1 到 256,每个宽度上都是 0.572,因为那个一直在长的尺度被除掉了。改成除以 d 会矫枉过正,把离散度压到 1/√d,每一行都被抹平成均匀平均。本文这个玩具头 d = 2,除数是 1.414; GPT-2 small 每个头 d = 64,除数正好是 8。
掩码,以及它通常是怎么被写错的
语言模型被训练去预测下一个 token。而注意力本身,并不阻止它先把那个 token 读一遍。
§03 那个方阵允许位置 0 去注意位置 5,而在训练时,那正是模型被要求猜的那个 token。所以解码器会在 softmax 之前把对角线以上的一切删掉。一行一行地取走,看未来怎么消失:
虚线方块是算了又被扔掉的分数 ——,这也是为什么一个没有融合的实现要做 36 次点积,只为留下 21 个。掩码是固定的、没有参数的,而且是整个机制里唯一知道“位置”这回事的部分。
活下来的部分会被重新归一:这一行仍然必须加起来等于一,所以剩下的 token 就把整份混合分掉。让提问的位置沿着句子往下走:
在 处,混合是 the 上的 1.00 —— 第一个 token 只有一个候选,所以不管它的查询说什么,它的输出就是它自己的值向量。到,权重摊在五个上,cat 拿 0.34。
现在说那个错误。0/1 掩码是个矩阵,拿它相乘看上去就是最自然的施加方式。把控件从 + (−∞) 切到 × 0:
因为被掩住的分数变成了 0 而不是 −∞,而 exp(0) = 1,ran 在一份它本不该出现的混合里仍然占着 0.092 —— 模型正在偷看答案。没有异常抛出,形状全对,训练损失还降得比平时更快,而这恰恰就是马脚。正确的写法是 s = s.masked_fill(m, float("-inf")),并且值得加一条断言:权重矩阵被掩住的位置必须严格为零。
那个加权和,以及它的账单
权重从来就不是答案。它们只是把值向量倒在一起时的配比。
最后一步只有一行:把每个值向量乘上它的权重,六个加起来。因为每个权重都是一的一部分,每一项都是沿那个值方向迈的一小步。一项一项地加上去:
注意这条折线从不折返,因为没有负权重能把它拉回去。六步之后,就是 it 的新表示:[0.82, 0.81] —— 一个已经被告知了一点关于那只猫的信息的 token。
这也是 §04 那条断言的几何形态。权重为正且加起来等于一,意味着输出是一个凸组合,于是它必然落在值张成的那个多边形里面。随便怎么拖 cat 的分数,试试能不能把它拽出去:
出不去。,输出停在 [1.21, 0.93],比cat 自己的值差了百分之一 —— §04 那条严格为正的下界,正是它能逼近一个角却永远到不了的原因。一个注意力头永远只能返回句子里已有内容的一个平均 —— 这正是它后面要接一个非线性 MLP 的原因,也是为什么一叠没有 MLP 的注意力层会塌缩成接近单个线性映射的东西。
接下来说这个机制做不到的事。这个和是在一个集合上求的,所以重排 token 只会重排权重,别的什么都不会变。打乱阅读顺序,盯着输出看:
柱子重排了,而 [0.82, 0.81] 一动不动。自注意力是置换等变的:在没有位置信息时,它分不出 the cat ate and it ran 和它的任何一种打乱,而且它是不声不响地失败 —— 模型照训,损失照降,词序压根就没进过表示。位置编码存在的意义就是打破这个对称性,而上一节那个因果掩码,是本页唯一另一个知道顺序的东西。
账单是:每一个有序对做一次点积,做两遍 —— 一遍给 Q · Kᵀ,一遍给那个混合。也就是 2n²d 次乘加;如果那个矩阵被真的写出来,还要 n² 个分数的内存,而 token 自己只占 n · d。拖动上下文长度:
在 处,一个头的分数矩阵在 fp16 下是 2.0 MiB —— 12 层 12 头合起来 288 MiB。到 时它是 32.0 GiB,而它所依据的那些 token 只有 192.0 MiB:小 171 倍。 FlashAttention 存在的理由,就是永远不把这些字节写下来。
四行代码,和它们画出来的东西
本页每一个想法,都落在这四行里。
s = q @ k.transpose(-2, -1) / d_k**0.5
s = s.masked_fill(mask, float("-inf"))
w = s.softmax(dim=-1)
out = w @ v把它们跑遍整句话,结果就是一张图:六行权重,每一行是一个查询在键上的混合,未来已被切掉。把掩码关掉,看右上三角回来:
这就是每篇论文都在印的那张注意力图,而它现在可以一格一格地读:因果掩码下 it 从 cat 取走 0.34 的混合,不掩码时是 0.30。四种能把它弄坏、又都不会抛异常的写法:掩码用乘不用加、没有位置信息、漏掉 √d 除数,以及键没有归一化。