3.7 注意力机制
1.把所有东西压进一个状态,太挤了
说起来,上一章我们留了个尾巴,循环网络要把整段历史塞进一个固定大小的隐藏状态里。可序列一长,这个状态就装不下了,早期的内容很容易被后面新来的内容盖掉。这就好比一个塞得太满的抽屉,新东西一塞进去,老东西就被挤到看不见的角落里,再想翻出来可就难了。
注意力机制(attention)换了个很不一样的思路。模型在生成当前这一步的表示时,可以直接回头查看之前每一个位置留下的表示,想看哪个就看哪个,看多久、看多深,都由模型自己根据当下的需要来决定。这就好比做开卷考试,你不用把所有知识点死背在脑子里,遇到题目随时翻回前面相关的段落查一查就行,一下子从容多了。我之前看过一本讲图书馆学的书,里头提到一个说法,好的检索系统讲究给每份资料贴好标签、随用随取,不必把所有知识一股脑塞进同一个抽屉里,这个道理放到注意力机制上我觉得格外贴切。
2.三个角色:查询、键和值
注意力机制借用了资料检索的思路,给每个位置都安排了三个角色。
第一个叫查询(query),记作 ,表示当前这一步想找什么样的信息。第二个叫键(key),记作 ,表示第 个候选位置身上贴的标签,专门用来被查询比对(下标 表示第几个位置)。第三个叫值(value),记作 ,表示这个候选位置真正能提供出来的内容。
打个找书的比方。你心里想着要一本讲深度学习的书(这就是查询 ),走到书架前扫一眼每本书的书名和标签(这些就是键 ),书名和你的需求越对口,你就越想伸手去拿那本,最后你真正读到的内容是书里的正文(这就是值 )。
不妨再换个行业的场景看看。小张在一家投资公司做行业研究,他这天想找的是近三年关于新能源车销量的一手数据(查询 ),于是他在资料库里一条条扫过去,每份报告的标题、关键词和发布时间就是键 ,哪份报告的标签跟他的需求对得上,他就把那份报告点开,真正读到的图表和结论就是值 。医生看病其实也走的是同一套路数,病人的主诉症状是查询,各项检查指标是键,最后给出的诊断依据是值,各行各业的检索逻辑骨子里都长得差不多。
具体怎么算呢?查询 跟每一个键 做点积(dot product,说白了就是两个向量对应位置相乘再相加,结果越大说明两者越像),得到的那个数,就表示这对查询和键有多相似。对所有位置都算完之后,模型用Softmax这一步,把这一堆相似度变成一组加起来等于 的权重 ( 表示第 个位置分到的注意力比重)。最后,把各个位置的值 按这个权重加起来,得到输出:
哪个位置的权重越大,它对最终输出的影响就越明显。模型就是这么通过权重,把注意力分配给了最相关的几个位置。
Softmax这一步到底干了什么,值得单独掰开揉碎说说。假设一共有3个位置,点积算出来是3个数,我们记作 、、( 就是查询和第 个键的点积结果)。Softmax分两步走。第一步给每个 套上一层指数函数,得到 ,这里的 是自然底数,约等于 ,这一步能保证所有数都变成正的,而且原来数值大的经过指数放大之后会更显眼。第二步把全部 加起来当分母,让各自去除,写成公式就是 ,分母里的 表示把每个位置 的 全部累加起来, 是求和符号。这么一除,每个 都落在 到 之间,而且全部加起来正好等于 ,这就成了一组正经的比重,能直接拿来当权重用。
光看公式可能还是有点悬,我们再换个几何角度看加权和这件事。把每个值向量 想成平面上从原点出发的一个箭头,权重 想成拉力的大小。加权和 就相当于所有箭头一起发力,最后合力指向的那个位置就是输出 。哪个位置权重越大,输出就被拉得越偏向那个位置的值向量。权重加起来等于 还有另一个好处,输出始终落在所有值向量围起来的那片范围里头,不会一下子飞到没边的地方去,这跟开卷考试翻书翻得再狠也翻不出书架的范围是一个道理。
我们再拿一组带数字的小例子从头手动走一遍,这样最踏实。设查询 ,这是一个二维向量,第一个分量是 、第二个分量是 。手头一共有3个位置,它们的键和值分别是 对应 , 对应 , 对应 。先算点积,,,。可以看出第1个和第3个位置跟查询最对口,第2个位置完全不对口。接着套Softmax,,,,分母就是 。三个权重分别是 ,,,三个加起来正好是 。最后做加权和,,第一个分量是 ,第二个分量是 ,所以输出就是 。你看输出明显偏向第1个和第3个位置的值向量,跟它们权重更大完全对得上,这就是注意力机制干活的全过程。
3.缩放点积注意力:批量化还要稳一手
实际用的时候,查询和键都是一整批一整批一起算的,一个个来根本不现实。查询和键的维度记作 ( 就是查询或键向量里头有几个数)。把所有查询拼成矩阵 ,所有键拼成矩阵 ,所有值拼成矩阵 ,整个注意力计算可以一次写成:
这里头的 是用矩阵形式,一次性算出所有查询和所有键的点积, 表示转置。分母里那个 是用来给点积的结果做一次缩放的。为什么要缩放呢?因为维度 越大,点积算出来的数值就越大,大到一定程度,Softmax的输出会变得几乎只集中在一个位置上(其它位置权重全被压成接近零),这一集中,梯度就不好传了。除以 能把数值压回到一个合理范围,让权重的分配别那么极端,训练也就稳了。
4.自注意力和交叉注意力:看自己还是看别人
到这里有个关键问题一直没说清楚,查询、键、值这三个角色到底是从哪儿冒出来的?答案得分两种情况来看。
第一种叫自注意力(self-attention),查询、键、值全都来自同一个输入序列。换句话说,序列里每个位置都在跟自己同一批兄弟比对,看看队伍里还有谁跟自己相关。举个具体场景,做情感分析处理一条商品评论的时候,评论前半段可能在夸味道不错,后半段突然开始吐槽服务态度差,模型要是光盯着前半段会以为这条评论在表扬,得让后半段那些吐槽的词和前半段夸人的词直接互相参考一下,模型才能摸清这条评论到底是在真心夸还是在阴阳怪气。自注意力干的就是让句子里每个词都能去打量其他词、互相印证的活儿。
第二种叫交叉注意力(cross-attention),查询来自一个序列,键和值则来自另一个序列。这就意味着一个序列在向另一个序列发问,去借自己需要的信息。最典型的场景是机器翻译。解码器正在一个词一个词地往外蹦英文的时候,每蹦一个词,查询 都是从英文这边出发的,而键 和值 全部来自编码器那边对整句中文编码出来的表示。英文每生成一个词,都会回头去中文那边按需取用对应的信息,正是交叉注意力在两种语言之间牵线搭桥。
5.注意力到底强在哪里
注意力本质上是一种按内容来找信息的办法。模型可以根据当前正在处理的内容,动态决定去参考序列里的哪些位置,而且这种参考是直接点对点的,不用像循环网络那样隔着一堆时间步慢慢传。
这就带来一个特别大的好处,哪怕是序列里隔得很远的两个词,注意力也能让它们直接发生联系,长距离依赖一下子变得好处理多了。我们回头细看一眼 3.5 节那个寝室传话游戏的窘境。循环网络里第1个词的信息想影响到第20个词,得老老实实经过中间那一大串隐藏状态的接力,每传一步都有可能被新进来的内容稀释一点、盖掉一点,传到后面早就面目全非了,这正是长距离依赖搞不定的根源。注意力机制干脆给每个位置都发了一个对讲机,第20个词想用第1个词的信息,一步直接调取就行,中间那些传递环节全省了,信息几乎不失真。
说穿了,这件事不只是模型才有的烦恼。记得有一次,小张跟我吐槽,他们组里有个项目,前面负责采集数据的同事把一条关键说明写在很早的文档里,后面接手的人一个传一个,传到最后写代码那位同事手里,意思已经完全变样了,整个项目返工了好几天。这其实就是工程版的传话游戏,注意力机制要解决的恰好就是这一类麻烦,让后面的人随时能直接翻回最早的原文核对,不必依赖中间一长串转述。
也正因为它这么管用,注意力成了下一章主角Transformer的地基,现代大模型那一身本事,相当一部分都是从这儿长出来的。
按惯例这里该收个尾,今天就先到这儿,下一章见。
练习
Q1. 注意力机制里的 query、key、value 三个角色分别是什么意思,又怎么从它们一步步算出输出?
query 是当前这一步想找什么样的信息,key 是每个候选位置身上用来被比对的标签,value 是候选位置真正能提供的内容。计算时先用 query 跟每个 key 做点积得到相似度,再过 Softmax 把这些相似度归一成一组加起来等于 1 的权重 ,最后把各位置的 value 按这个权重加权和 ,就是输出。
Q2. 沿用本章那个手算例子:,三个位置的 ,,。请算出三个点积值,并说明第 2 个位置为什么权重最小。
,,。三个 Softmax 权重约为 、、,加起来正好是 1。第 2 个位置点积为 0,是三者里最小的,说明它的 key 跟 query 完全不对口,所以分到的权重最小。
Q3. 缩放点积注意力里那个分母 是多余的装饰吗,去掉会怎样?
不是。维度 越大,点积的数值就越大,大到一定程度 Softmax 的输出会几乎集中在一个位置上(其余权重被压成接近 0)。这一集中梯度就近乎为零,模型学不动。除以 把点积数值压回温和范围,让权重分配不那么极端,训练才稳。
Q4.(面试题) 自注意力和交叉注意力有什么区别?注意力机制相比循环网络,在处理长距离依赖上强在哪里?
自注意力里 query、key、value 都来自同一个输入序列,序列内每个位置都跟自己这批兄弟互相比对;交叉注意力则是 query 来自一个序列、key 和 value 来自另一个序列,比如机器翻译里解码器生成的每个词(query)回头去编码器那边(key、value)取用信息。相比循环网络,注意力的信息传递是直接点对点的:第 20 个词想用第 1 个词,一步直接调取就行,不用隔着一堆时间步慢慢接力,所以长距离依赖信息几乎不失真,这正是它能甩开循环网络的地方。