3.8 Transformer与表示学习

1.一场改写历史的小组会

时间回到2017年。那时候你要是想做机器翻译,比如把一句英语译成中文,最顶配的方案是循环神经网络配上注意力,行话叫seq2seq加注意力。这套方案翻译质量确实不错,可它身上有个绕不开的毛病。循环网络得一个词一个词往后读,前一个词处理完才轮到下一个,整条流程是串行的。

串行是什么概念呢,不妨类比成食堂二楼那个永远只开一个窗口的兰州拉面档口。你再急,前面那位还在纠结加不加香菜、要不要加肉,你也只能在后头干等。GPU这种算力怪兽最擅长的本是一口气算一大片,可循环网络这种一个接一个的串行打法根本喂不饱它,训练起来慢得让人心焦,别人都把期末卷刷完了,你这边一道题还没跑完。

当年在谷歌,有一组研究员专门盯着这个痛点,带头署名的是位叫Vaswani的年轻研究员。他们琢磨出一个大胆到有点离谱的念头:循环网络又慢又难训,干脆把它整个扔掉,只留注意力,行不行?论文标题就透着这股劲,《Attention Is All You Need》,注意力就是你需要的全部(这标题还顺手借了甲壳虫乐队那首《All You Need Is Love》的梗)。这篇论文提出的全新模型结构,就是后来横扫整个领域、被推上神坛的Transformer。

这群研究员自己当时多半也没料到,这篇本来只是冲着翻译去的论文,会像一颗深水炸弹,把整个领域搅得天翻地覆。Transformer不光翻译做得又快又好,更关键的是它那种结构特别适合往大了堆,直接铺开了通向后来BERT、GPT乃至ChatGPT的路。说起来,今天大模型这一波泼天富贵,源头都扎在2017年这场小组会上。

2.先说最核心的那一招:自注意力

上一章我们把注意力的三兄弟query、key、value拆解过了,还记得吧。Transformer的灵魂,就是把这招用到了一种叫自注意力(self-attention)的玩法上。这一节我们从最简单的样子讲起,因为它是后面所有东西的地基,地基没打牢,上面的多头、残差、前馈全是空中楼阁。

先快速回忆下三兄弟。query(查询)是当前位置想找什么样的信息,key(键)是每个候选位置身上贴的、专门用来被比对的标签,value(值)是候选位置真正能拿出来的内容。第30章里我们记作 qqkjk_jvjv_j,这里照旧。

自注意力妙就妙在一个自字。在Transformer里,序列里的每一个位置,都会从自己那条输入向量,同时变出query、key、value三样东西。设第 ii 个位置的输入向量是 xix_iii 表示第几个位置),它分别通过三个权重矩阵 WQW_QWKW_KWVW_V(这三个矩阵都是模型在训练里学出来的参数),变出三样东西:查询 qi=WQxiq_i=W_Qx_i、键 ki=WKxik_i=W_Kx_i、值 vi=WVxiv_i=W_Vx_i

也就是说,每个位置都自带query、key、value三件套。接下来模型让每个位置的query,去和所有位置的key比对相似度,再拿这个相似度当权重,把所有位置的value加权求和,聚合出来的结果就是这个位置的新表示。把这一步写成矩阵形式,就是上一章那个缩放点积注意力:

Attention(Q,K,V)=softmax(QKTdk)V\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^\mathsf{T}}{\sqrt{d_k}}\right)V

这里的 QQKKVV 分别是把序列里所有位置的query、key、value拼起来的三个大矩阵,QKTQK^\mathsf{T} 一次性算出所有位置两两之间的相似度,T\mathsf{T} 是转置,dkd_k 是key向量的维度,除以 dk\sqrt{d_k} 是为了防止维度太大时点积数值爆掉、把softmax压成只盯着一个位置。这公式第30章细讲过,这里就不重复展开了。

为什么偏偏除以 dk\sqrt{d_k},这个数怎么来的。 这值得推一下,不然总觉得是拍脑袋。假设 qqkk 的每个分量都是独立的、均值0、方差1的随机变量,那么点积 qk=i=1dkqikiq\cdot k=\sum_{i=1}^{d_k}q_i k_i 的方差是(独立变量乘积的方差可加):

Var(qk)=i=1dkVar(qiki)=i=1dkVar(qi)Var(ki)=dk\operatorname{Var}(q\cdot k)=\sum_{i=1}^{d_k}\operatorname{Var}(q_i k_i)=\sum_{i=1}^{d_k}\operatorname{Var}(q_i)\operatorname{Var}(k_i)=d_k

也就是说点积的标准差是 dk\sqrt{d_k},它随维度 dkd_k 增长。维度一大,点积数值动辄正负好几十,softmax 一算全是 0011 的极端值(梯度几乎为零,学不动)。除以 dk\sqrt{d_k} 正好把方差拉回 11(标准差变回 dk/dk=1\sqrt{d_k}/\sqrt{d_k}=1),让 softmax 的输入落在一个温和的范围内,梯度健康。这就是 dk\sqrt{d_k} 的由来——它精确地抵消了点积方差随维度增长的部分

只要想清楚一件事就通了:因为query、key、value全都来自同一句话自己,所以一个词可以直接去看句子里任何一个词,包括它自己。这就好比做开卷考试,不仅能翻前面,整张卷子哪里都能翻。比如我爱你的我,可以直接去看一眼你长什么样,从而明白这两个词之间有联系。这种全局互看的能力,就是自注意力最值钱的地方,也是Transformer能甩开循环网络的根本原因。

3.多头注意力:一组人各看各的角度

刚讲的是一个注意力头的情况,它一次只能在一个表示空间里找关系,能力多少有点单一。多头注意力(Multi-Head Attention)就是专门来破这个限制的。

它的做法是把表示拆到好几个子空间里:把原来一整块的 QQKKVV 切成 hh 份(hh 就是头的个数),每一个头拿自己那一份小 QQKKVV,独立地算一次自注意力,算完把 hh 个头的结果拼到一起,再过一个线性变换送进下一层。这么一来,不同的头就能在不同的子空间里各找各的关系,互不打架。

这件事不妨类比成大学时几个同学组队啃一篇paper,一份论文发下来,舍长看公式推导、有人查作者机构、有人翻引用列表顺藤摸瓜、有人盯实验图表,大家各看各的角度,最后把线索汇到一起,理解自然比一个人闷头啃全面得多。训练完打开模型一看,不同的头还真会自发地学会关注不同类型的关系,有的头专门看近距离的词语搭配,有的头专门看隔得很远的呼应,分工明确得很,跟组队啃论文是一个道理。

4.把零件拼成一整层

光有自注意力还拼不成一个Transformer,得把它和另外几样零件拧成一整层,再把层一层一层摞起来,才算完事。

一个标准的Transformer层是这么一条流水线。输入先进自注意力子层,让不同位置互相交流信息,出来之后,把它和这一子层的输入相加(这种把输出和输入相加的操作叫残差连接,第27章讲ResNet的时候提过),相加完再做一次归一化(把数值范围拉回一个稳定的分布)。接着进前馈网络子层,对每个位置单独来一次非线性变换,出来之后再次做残差相加、再归一化。一层就这么走完一轮。

这里头的前馈网络、残差连接、归一化,前面章节都讲过,这一层没什么新零件,只是把自注意力跟这几样拧到一起,让它既能交流信息、又能加工信息、还能稳稳地训。把很多这样的层摞起来,就是完整的Transformer。每过一层,表示就被重新整理一遍,越往后越精。

哦对了,还有个历史小知识得补一句。原始论文里的Transformer是为翻译设计的,所以分成两半,一半叫编码器(encoder),负责把原句吃透,另一半叫解码器(decoder),负责一句一句把译文吐出来。后来大家发现这两半可以拆开单独用:BERT只用了编码器那一半,所以特别擅长理解语言,GPT只用了解码器那一半,所以更擅长接着往下生成文字。今天你天天拿来问作业、润色文案甚至帮你写实验报告的那些聊天大模型,基本都是GPT那条只用解码器路线的徒子徒孙。

5.位置表示:注意力本来是不认顺序的

这里有个大坑得专门拎出来讲一下。注意力本身干的事,说穿了就是比较两个位置的内容像不像,它压根不在乎谁先谁后。可是在语言里,顺序这东西要命地重要。我爱你,三个字原样念是一种意思,倒过来念成你爱我,意思天差地别。再比如上课点名,第一个到的人和踩着铃声冲进来的人,老师心里的印象完全两码事。所以模型必须额外把位置信息给补回去,不然注意力就抓瞎了。

补位置信息这件事,由位置编码来干。它的套路是这样的,把每个位置的编号变成一个向量,再把这个位置向量加到对应词元的表示上。加了这一笔,模型就能区分两个内容相同的词出现在句子不同位置时,其实扮演的是不同角色。你想想看,要是没这个位置编码,Transformer看一句打乱顺序的话和看原句,会觉得一模一样,那还翻译什么呢。

后来位置编码也演化出好几种花样。原始论文里用的是固定的一套正弦余弦公式,不用学、直接给定。这套公式长这样: 对位置 pospospospos 是从0开始的位置编号)和维度索引 iiii 是位置向量里的第几维),偶数维用正弦、奇数维用余弦:

PEpos,2i=sin ⁣(pos100002i/d),PEpos,2i+1=cos ⁣(pos100002i/d)PE_{pos,2i}=\sin\!\left(\frac{pos}{10000^{2i/d}}\right),\qquad PE_{pos,2i+1}=\cos\!\left(\frac{pos}{10000^{2i/d}}\right)

这里 dd 是位置向量的总维度,2i2i2i+12i+1 分别表示偶数维和奇数维的下标,1000010000 是个固定的大数。不同维度用不同频率的正弦波——低维(ii 小)频率高、变化快,高维(ii 大)频率低、变化慢,整个位置向量就像一组不同刻度的钟表,组合起来能唯一地编码每一个位置。

为什么用正弦余弦能表达相对位置。 这套公式有个巧妙的性质:对任意固定的偏移 Δ\Delta,位置 pos+Δpos+\Delta 的编码可以表示成位置 pospos 编码的线性函数。拿一对相邻的正弦余弦维度看,用三角恒等式展开:

(sin(ω(pos+Δ))cos(ω(pos+Δ)))=(cos(ωΔ)sin(ωΔ)sin(ωΔ)cos(ωΔ))(sin(ωpos)cos(ωpos))\begin{pmatrix}\sin(\omega(pos+\Delta))\\\cos(\omega(pos+\Delta))\end{pmatrix}=\begin{pmatrix}\cos(\omega\Delta)&\sin(\omega\Delta)\\-\sin(\omega\Delta)&\cos(\omega\Delta)\end{pmatrix}\begin{pmatrix}\sin(\omega pos)\\\cos(\omega pos)\end{pmatrix}

ω=1/100002i/d\omega=1/10000^{2i/d} 是该维度的频率。)那个 2×22\times2 矩阵只依赖偏移 Δ\Delta,不依赖绝对位置 pospos。这意味着模型只要学到"相对偏移 Δ\Delta 对应的线性变换",就能从任意位置的编码推算出偏移 Δ\Delta 后的位置编码——这就是相对位置信息被编码进去的数学原因。后人又做出了可以学习的那种位置编码,让模型在训练里自己慢慢调出一套最合适的位置向量。不管哪一种,目的都一样,把顺序这条命脉给续上。

6.为什么Transformer能赢:能并行,还能往大了堆

Transformer能一统江湖,靠的是两把刷子,而且这两把都跟它的并行能力脱不开干系。

第一把刷子是能并行。循环网络得一个词一个词串着往后算,前面没算完后面就干等着没法动。Transformer的自注意力可以一口气把序列里所有位置同时算出来,谁也不用等谁。这下GPU那种一口气算一大片的并行算力终于被喂饱了,训练速度直接起飞,原来训一轮的工夫现在能训上好几轮。记得小张第一次跑Transformer翻译任务,原先用循环网络训一轮要守一整晚,换成Transformer之后几个钟头就跑完一轮,他感叹总算不用熬夜盯着显卡了。

第二把刷子,也是更要命的一把,是能往大了堆。正因为能并行、训得快,人们才敢放手把Transformer越摞越大,参数从几千万一路堆到几十亿、几千亿。它还特别争气,越大越聪明,几乎没碰着明显的天花板,Scaling Law那套规律在这一路简直势如破竹。我之前翻过一本讲半导体发展史的书,里头提到摩尔定律几十年间把芯片上的晶体管数推高了成千上万倍,Transformer这条规模之路其实有点那个味道,能堆就尽量堆。这一下,规模这条路就被彻底打开了。BERT、GPT、ChatGPT,这些后来响当当、把整个AI赛道带上新高度的名字,地基清一色全是Transformer。

7.表示是一层层学出来的

Transformer的每一层,其实都在重新整理输入的表示。底层还保留着比较多局部和词形方面的信息,比如这个词长什么样、和旁边几个词怎么搭配。中间层开始把更远的上下文揉进来,整句话的意思慢慢成形。到了高层,表示就越来越接近任务真正需要的抽象关系,谁跟谁搭、谁在否定谁、整段话到底在讲什么。层数堆得越多,模型的表达能力通常就越强,但同时它也更需要海量数据来喂饱,对训练稳定性的要求也更高,一个不留神就摆烂给你看。

理解了这一点,你就能明白为什么现在的大模型动不动就几十上百层,还得配上几万亿词的训练语料。这背后那股子把Transformer往极限去堆的执念,全都是从2017年那篇论文开始的。说Transformer是这十年最重要的发明之一,真一点都不过分。

练习

Q1. Transformer 为什么要单独加位置编码,光用自注意力不行吗?

不行。自注意力本质只是比较两个位置的内容像不像,它对位置顺序没有概念,把句子打乱顺序和看原句会觉得一模一样。可语言里顺序要命地重要("我爱你"和"你爱我"意思天差地别),所以必须额外用位置编码把顺序信息补回去,模型才能区分内容相同但出现在不同位置的词。

Q2. 为什么缩放点积注意力里要除以 dk\sqrt{d_k}?请从方差角度说明。

假设 qqkk 的每个分量都是均值 0、方差 1 的独立随机变量,那么点积 qk=i=1dkqikiq\cdot k=\sum_{i=1}^{d_k}q_ik_i 的方差是 dkd_k(各分量乘积方差可加),标准差是 dk\sqrt{d_k},它随维度增长。维度一大点积数值动辄正负好几十,Softmax 一算全是 0、1 的极端值,梯度几乎为零。除以 dk\sqrt{d_k} 正好把方差拉回 1,让 Softmax 输入落在温和范围,梯度才健康。这恰好抵消了点积方差随维度增长的部分。

Q3. 一个标准 Transformer 层的流水线是怎么走的?多头注意力又是干嘛的?

一层流水线是:输入先进自注意力子层(不同位置互看),输出和子层输入做残差相加、再归一化;接着进前馈网络子层做逐位置非线性变换,出来再残差相加、再归一化,一层才走完。多头注意力是把表示切到 hh 个子空间,每个头拿自己那一份小 QQKKVV 独立算一次自注意力,算完拼起来再过线性变换,让不同头在不同子空间各找各的关系(如近距离搭配、远距离呼应),理解更全面。

Q4.(面试题) Transformer 为什么能赢过循环网络一统天下?再说说 BERT 和 GPT 各自用了 Transformer 的哪一半,分别擅长什么。

靠两把刷子,都跟并行能力有关。一是能并行:循环网络得一个词接一个词串行算,自注意力可以一口气把所有位置同时算出来,喂饱 GPU 的并行算力,训练速度起飞。二是能往大了堆:正因为训得快,人们才敢把参数从几千万堆到几千亿,而且越大越聪明几乎没碰天花板。BERT 只用了编码器那一半,擅长理解语言;GPT 只用了解码器那一半,更擅长接着往下生成文字,今天的主流聊天大模型基本都是 GPT 那条解码器路线的徒子徒孙。

相关标签
深度学习Transformer注意力PyTorch