多头注意力 入门
自注意力给每个查询一行 softmax —— 一份混合。真实句子对一个 token 的要求,比一份混合能诚实回答的更多。多头注意力的全部把戏,用你自己能拖动的图逐步拆开:把同一条宽度切成 h 份,独立跑 h 次完整的注意力计算,让它们像真实训练出来的头那样各自专精,再拼回去 —— 花的参数和算力,跟原来一个宽头一模一样。
一行,一份混合
自注意力给每个查询恰好一行 softmax。这一页要讲的,是一行不够用的那一刻。
一个查询的输出,是值向量的加权平均,权重来自一次 softmax。这不是实现细节 —— 这就是一个头能产出的东西的全部形状:一份混合,不管权重在六个键之间怎么分。
还是看 it。给它配一个头,这个头只会盯着紧挨在前面的那个 token,把它的查询拖到句子里任何位置:
注意这个头对语法毫无想法 —— 它根本不知道 it 是个代词,只知道 and 在它前一个位置。这是一个真实、有用的信号,但也不是 it 唯一需要的那个。
再配一个头,让它去盯查询在语法上依附的那个词。看着两行随查询一起移动 —— 到时,两个头指向的完全是两个不同的 token:
因为 it 是 ran 的主语,第二个头够到的是前方 —— 和第一个头一样,只差一格,方向却相反:第一个头的规则只能往后看,永远表达不出向前看。一行,两个都真实、都不同的答案。
一个头没有第三种选择:它只能混合。把混合比例从纯位置滑到纯语法,读一下经过中点时的峰值:
看着峰值权重在经过中点时往下掉,再往另一个纯答案那边爬回去 —— 中间这一段,这一行从来没有真正靠近过任何一个 token。,它谁都没靠近。
这个凹陷不是这句话独有的巧合。把它对着混合比例画出来,这个头在两个信号上能拿到的最高权重,恰恰在它最努力想兼顾两边的地方最低:
妥协不是第三个答案 —— 它是那两个答案各自的一份更差的拷贝。下一页的解法,不是把一个头变聪明,而是不再让一个头身兼两职。
同一条宽度,切成 h 个头
两个头意味着两行 softmax。不意味着 token 变宽了。
token 到达时仍然是一条向量,宽度是 d_model 个数字 —— 下面的图里是 8,GPT-2 small 里是 768。多头注意力不会往这条宽度上加东西,它只是把这条宽度切成 h 份等宽的小段,每段宽 d_k = d_model / h。
拖动改变 h,看同一条宽度被切出不同的样子 —— 根本没有切口,一个头,整条向量:
看着每一段随着每加一刀变窄,而两端从不移动。四个宽度为 2 的头,逐维度算下来,和一个宽度为 8 的头是同一份预算 —— 只是换了一种排法。
这里最容易招来的一个误会:很容易把每个头想象成在读 token 自己的四分之一。在两种画法之间切换,看看每个头在输入端到底看到了什么,而不是输出端:
因为每个头的投影都是它自己完整的 d_model × d_k 矩阵,每个头都读到了全部八个输入数字,只是留下一段较窄的结果—— token 本身从没被切开过。
真正被切开的是 W_Q 自己的输出列—— 以及 W_K、W_V 的输出列,同样的切法,各自独立。再改一次 h,看这三个一起被切开:
查询、键、值各自守着自己的 d_model × d_model 矩阵、各自的切法 —— 它们之间唯一共享的只有数字 h。第 3 节要讲的,是这些窄条各自单干时会做什么。
不再混合 —— h 次完整的计算
每条窄切片都独立跑一遍完整的注意力机制:自己的分数,自己的 softmax,自己的加权和。
一个 d_k = 2 的查询和键组成的头,不是注意力的简化版 —— 它就是注意力本身,只是宽度更窄。打分、softmax、加权和,没有一步会因为这是多个头之一而变样。
拖动查询,看跟着位置走的那个和跟着语法走的那个同时以全力给出答案,中间没有任何东西挡在它们和自己那一行之间:
看着这两行,谁也没有向对方软化。第 1 节里的那个妥协不见了 —— 不是被修好了,是被替换掉了。两份独立的混合,花费和当初那一行被迫混合的成本一样。
这不会只停在两个。把第 1 节没机会用上的两个头也加进来 —— 一个不管查询是谁,永远拉向同一个 token,第四个则谁都不肯全力押注:
四行里有三行有一个清晰、能叫得出名字的故事。第四行没有,正因如此这里没给它上色 —— 第 4 节要讲的,就是在一个真实训练过的模型里,这四种到底意味着什么。
独立性还咬人的另一个地方:每个头除以的是自己的√d_k,不是整个模型的。,每个头都会悄悄变平,随着 d_k 变小:
因为正确的除数恰好把宽度约掉了,同样强度的信号的峰值权重在任何 d_k 上都稳定不变。错误的那个除数会把这份信号冲淡 —— 不报错,不崩溃,只是 softmax 永远没机会自信起来。
真实的头,最后长成了什么样子
第 1—3 节手工搭出来的那两种头,不是为了方便讲故事编出来的。真实训练出来的头,会自己分化成能叫得出名字的类型。
Voita、Talbot、Moiseev、Sennrich 和 Titov 训练了一个翻译模型,看它的头到底在注意什么。大多数头落进三种角色之一 —— 本页已经出现的这两种,再加第三种:不管查询是谁,都拉向句子里最罕见的那个 token。
把跟着位置走的头和永远伸向最罕见 token 的头放在一起,把查询拖着走遍整句话:
看着上面那行的峰值每次都跟着查询晚一步走,下面那行则一步都不挪。同样的六个键,同样的控件,两条结构上完全不同的规则。
把每个查询位置都画出来,这个差异就变成了一个形状:位置头描出一条笔直的对角线,稀有词头是一条钉死在 cat 上的水平线:
句法头那条线两者都不是 —— 它在每个位置按依存关系实际指向哪里就跳到哪里,在 it 处往前跳一格,在 ate 处往后退一格。单独走一遍,看这个目标怎么都不肯落进一条规律里:
因为这个头回答的是「这依附于谁」,不是「往回数几格」,它的目标跟着语法走到哪算哪。这正是第 1 节里那一行单靠位置模仿不出来的信号。
回报在这:故事讲得这么干净的头,正是模型丢不起的那些。在 Voita 等人的翻译模型里,48 个头里,活下来的正是这几种专职头:
留下来的头绝大多数是位置头、句法头或稀有词头—— 第 3 节里那个说不清故事的第四种,恰恰是最先被剪掉的。专精的头挣得了自己的宽度;不专精的头,删掉它几乎不用付出代价。
四份输出,合回一份
每个头都各自算完了,上一层要的还是一条 d_model 宽的向量,不是四条窄的。
最直接的办法就是真正的办法:把四份输出并排放在一起。,宽度正好回到第 2 节切开它之前的那个数:
看着最后一个空位被填上,其他位置一动不动 —— 拼接不会碰任何一个已经放好的数,它只会把下一个头的两个值接到末尾。
这些不是占位符。对 token it 来说,每个头的输出都是对同样六个值向量做的一次真实加权和,只是权重是这个头自己的。滑到别的 token,这八个数每一个都是真的:
四个头,四个诚实的答案,首尾接在一起 —— 但到目前为止,它们仍然是四种独立的意见。没有谁在决定头一的输出时读过头三的数。
这正是 W_O 存在的理由。它又是一个 d_model × d_model 矩阵,一次性作用在整条拼接向量上。把它的混合程度从零往上调,看某一个位置怎么开始带上每个头答案的一份:
因为 W_O 不是分块对角矩阵,位置头自己的那个位置最终成了四者的混合 —— 这是整个机制里唯一一处,句法上的发现和位置上的发现被允许影响同一个输出数字。
到这里,四个矩阵已经把活干完了:W_Q、W_K、W_V 把 token 切成若干个头,W_O 再把这些头拼回去。四个矩阵是同一种形状:
这个共同的形状不是巧合 —— 它正是下一节要讲的全部论点。四个矩阵守着同一条宽度,不管这条宽度被切成了几个头。
h 个头到底要花多少
很容易以为 h 个头要花一个头 h 倍的成本。实际上分毫不差 —— 参数和算力都一样。
不管 h 是多少,W_Q、W_K、W_V、W_O 各自都是 d_model × d_model—— 第 5 节已经把这四个画成一样大小了。在这里拖动 h,看这个公式给出的参数量:
每一刀都改变 d_k,但从没碰过那个数字。四个矩阵守着同一条固定宽度,参数量的公式里根本没有 h 这一项 —— 不是大致不变,是精确不变。
注意力的矩阵乘法随 d 缩放,而 h 个宽度为d_model / h 的头加起来恰好是 d_model,所以总量是 h × d_model / h 次乘加 —— h 被约掉了。沿着 h 拖动,看这条曲线怎么都不肯离开底下那条参照线:
和四个窄头,在任何上下文长度上都画出同一条曲线—— 这是 GPT-2 small 自己的数字,不是四舍五入出来的示意图。四个投影矩阵再多花大约 1.5 倍,同样也不会挪动 —— 这一节讲的成本,没有一样跟这条宽度怎么切有关系。
但这不代表 h 是白拿的。把它调高,d_k 就会一直变窄 —— 同一根滑块,这次读的是一个头还能表达什么,而不是它花多少:
在 d_k = 1 时,查询和键都只是一个数字,同号的两个数字永远指着「同一个方向」——这么窄的头,能分辨的方向只剩两个,不管它想注意什么。不会报错。你刚才看到的 FLOPs 和参数量依然纹丝不动。这个头只是悄悄地没什么话可说了,这也是为什么已发布的模型都把 d_k 留在 64 到 128 之间,让 h 跟着 d_model 一起长,而不是超过它去长。
七行代码,买到了什么
本页每一个想法,都落在这七行里。
q = (x @ Wq).view(n, h, dk).transpose(0, 1) k = (x @ Wk).view(n, h, dk).transpose(0, 1) v = (x @ Wv).view(n, h, dk).transpose(0, 1) s = q @ k.transpose(-2, -1) / dk**0.5 w = s.softmax(dim=-1) out = (w @ v).transpose(0, 1).reshape(n, h * dk) out = out @ Wo
view 和 transpose 就是整个拆分 —— x 本身毫无变化,变的只是那三个投影自己的输出被怎么读。dk**0.5 而不是 d_model**0.5,是第 3 节的那个坑。最后 reshape 再 @ Wo,是第 5 节的拼接和混合。h 恰好出现四次,成本一次都没出现在这七行里 —— 同样这七行,不用改一个字,从 Transformer-base 的到 GPT-3 的,已发布模型选定的任何宽度都能跑起来:
d_k 挪动得比 h 小得多:三代模型里挪动最少的这个比例,是每个头能拿到的表达余量,不是模型的宽度。把一个宽头拆成若干个窄头,从来就不是为了多花点算力 —— 是为了不多花一分钱,就让每个 token 能拿到不止一个答案。这些头之间仍然说不清楚的是哪个 token 排在前面,而这正是下一页要开始讲的地方。