3.7 注意力机制

1.把所有东西压进一个状态,太挤了

说起来,上一章我们留了个尾巴,循环网络要把整段历史塞进一个固定大小的隐藏状态里。可序列一长,这个状态就装不下了,早期的内容很容易被后面新来的内容盖掉。这就好比一个塞得太满的抽屉,新东西一塞进去,老东西就被挤到看不见的角落里,再想翻出来可就难了。

注意力机制(attention)换了个很不一样的思路。模型在生成当前这一步的表示时,可以直接回头查看之前每一个位置留下的表示,想看哪个就看哪个,看多久、看多深,都由模型自己根据当下的需要来决定。这就好比做开卷考试,你不用把所有知识点死背在脑子里,遇到题目随时翻回前面相关的段落查一查就行,一下子从容多了。我之前看过一本讲图书馆学的书,里头提到一个说法,好的检索系统讲究给每份资料贴好标签、随用随取,不必把所有知识一股脑塞进同一个抽屉里,这个道理放到注意力机制上我觉得格外贴切。

2.三个角色:查询、键和值

注意力机制借用了资料检索的思路,给每个位置都安排了三个角色。

第一个叫查询(query),记作 qq,表示当前这一步想找什么样的信息。第二个叫键(key),记作 kjk_j,表示第 jj 个候选位置身上贴的标签,专门用来被查询比对(下标 jj 表示第几个位置)。第三个叫值(value),记作 vjv_j,表示这个候选位置真正能提供出来的内容。

打个找书的比方。你心里想着要一本讲深度学习的书(这就是查询 qq),走到书架前扫一眼每本书的书名和标签(这些就是键 kjk_j),书名和你的需求越对口,你就越想伸手去拿那本,最后你真正读到的内容是书里的正文(这就是值 vjv_j)。

不妨再换个行业的场景看看。小张在一家投资公司做行业研究,他这天想找的是近三年关于新能源车销量的一手数据(查询 qq),于是他在资料库里一条条扫过去,每份报告的标题、关键词和发布时间就是键 kjk_j,哪份报告的标签跟他的需求对得上,他就把那份报告点开,真正读到的图表和结论就是值 vjv_j。医生看病其实也走的是同一套路数,病人的主诉症状是查询,各项检查指标是键,最后给出的诊断依据是值,各行各业的检索逻辑骨子里都长得差不多。

具体怎么算呢?查询 qq 跟每一个键 kjk_j 做点积(dot product,说白了就是两个向量对应位置相乘再相加,结果越大说明两者越像),得到的那个数,就表示这对查询和键有多相似。对所有位置都算完之后,模型用Softmax这一步,把这一堆相似度变成一组加起来等于 11 的权重 αj\alpha_jαj\alpha_j 表示第 jj 个位置分到的注意力比重)。最后,把各个位置的值 vjv_j 按这个权重加起来,得到输出:

y=jαjvjy=\sum_j\alpha_jv_j

哪个位置的权重越大,它对最终输出的影响就越明显。模型就是这么通过权重,把注意力分配给了最相关的几个位置。

Softmax这一步到底干了什么,值得单独掰开揉碎说说。假设一共有3个位置,点积算出来是3个数,我们记作 s1s_1s2s_2s3s_3sjs_j 就是查询和第 jj 个键的点积结果)。Softmax分两步走。第一步给每个 sjs_j 套上一层指数函数,得到 esje^{s_j},这里的 ee 是自然底数,约等于 2.7182.718,这一步能保证所有数都变成正的,而且原来数值大的经过指数放大之后会更显眼。第二步把全部 esje^{s_j} 加起来当分母,让各自去除,写成公式就是 αj=esj/kesk\alpha_j=e^{s_j}/\sum_k e^{s_k},分母里的 kesk\sum_k e^{s_k} 表示把每个位置 kkeske^{s_k} 全部累加起来,\sum 是求和符号。这么一除,每个 αj\alpha_j 都落在 0011 之间,而且全部加起来正好等于 11,这就成了一组正经的比重,能直接拿来当权重用。

光看公式可能还是有点悬,我们再换个几何角度看加权和这件事。把每个值向量 vjv_j 想成平面上从原点出发的一个箭头,权重 αj\alpha_j 想成拉力的大小。加权和 jαjvj\sum_j\alpha_jv_j 就相当于所有箭头一起发力,最后合力指向的那个位置就是输出 yy。哪个位置权重越大,输出就被拉得越偏向那个位置的值向量。权重加起来等于 11 还有另一个好处,输出始终落在所有值向量围起来的那片范围里头,不会一下子飞到没边的地方去,这跟开卷考试翻书翻得再狠也翻不出书架的范围是一个道理。

我们再拿一组带数字的小例子从头手动走一遍,这样最踏实。设查询 q=[1,0]q=[1,0],这是一个二维向量,第一个分量是 11、第二个分量是 00。手头一共有3个位置,它们的键和值分别是 k1=[1,0]k_1=[1,0] 对应 v1=[1,0]v_1=[1,0]k2=[0,1]k_2=[0,1] 对应 v2=[0,1]v_2=[0,1]k3=[1,1]k_3=[1,1] 对应 v3=[1,1]v_3=[1,1]。先算点积,qk1=1×1+0×0=1q\cdot k_1=1\times1+0\times0=1qk2=1×0+0×1=0q\cdot k_2=1\times0+0\times1=0qk3=1×1+0×1=1q\cdot k_3=1\times1+0\times1=1。可以看出第1个和第3个位置跟查询最对口,第2个位置完全不对口。接着套Softmax,e12.718e^1\approx2.718e0=1e^0=1e12.718e^1\approx2.718,分母就是 2.718+1+2.718=6.4362.718+1+2.718=6.436。三个权重分别是 α1=2.718/6.4360.422\alpha_1=2.718/6.436\approx0.422α2=1/6.4360.155\alpha_2=1/6.436\approx0.155α3=2.718/6.4360.422\alpha_3=2.718/6.436\approx0.422,三个加起来正好是 11。最后做加权和,y=0.422×[1,0]+0.155×[0,1]+0.422×[1,1]y=0.422\times[1,0]+0.155\times[0,1]+0.422\times[1,1],第一个分量是 0.422+0+0.422=0.8440.422+0+0.422=0.844,第二个分量是 0+0.155+0.422=0.5770+0.155+0.422=0.577,所以输出就是 y[0.844,0.577]y\approx[0.844,0.577]。你看输出明显偏向第1个和第3个位置的值向量,跟它们权重更大完全对得上,这就是注意力机制干活的全过程。

3.缩放点积注意力:批量化还要稳一手

实际用的时候,查询和键都是一整批一整批一起算的,一个个来根本不现实。查询和键的维度记作 dddd 就是查询或键向量里头有几个数)。把所有查询拼成矩阵 QQ,所有键拼成矩阵 KK,所有值拼成矩阵 VV,整个注意力计算可以一次写成:

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

这里头的 QKTQK^\mathsf{T} 是用矩阵形式,一次性算出所有查询和所有键的点积,T\mathsf{T} 表示转置。分母里那个 d\sqrt{d} 是用来给点积的结果做一次缩放的。为什么要缩放呢?因为维度 dd 越大,点积算出来的数值就越大,大到一定程度,Softmax的输出会变得几乎只集中在一个位置上(其它位置权重全被压成接近零),这一集中,梯度就不好传了。除以 d\sqrt{d} 能把数值压回到一个合理范围,让权重的分配别那么极端,训练也就稳了。

4.自注意力和交叉注意力:看自己还是看别人

到这里有个关键问题一直没说清楚,查询、键、值这三个角色到底是从哪儿冒出来的?答案得分两种情况来看。

第一种叫自注意力(self-attention),查询、键、值全都来自同一个输入序列。换句话说,序列里每个位置都在跟自己同一批兄弟比对,看看队伍里还有谁跟自己相关。举个具体场景,做情感分析处理一条商品评论的时候,评论前半段可能在夸味道不错,后半段突然开始吐槽服务态度差,模型要是光盯着前半段会以为这条评论在表扬,得让后半段那些吐槽的词和前半段夸人的词直接互相参考一下,模型才能摸清这条评论到底是在真心夸还是在阴阳怪气。自注意力干的就是让句子里每个词都能去打量其他词、互相印证的活儿。

第二种叫交叉注意力(cross-attention),查询来自一个序列,键和值则来自另一个序列。这就意味着一个序列在向另一个序列发问,去借自己需要的信息。最典型的场景是机器翻译。解码器正在一个词一个词地往外蹦英文的时候,每蹦一个词,查询 qq 都是从英文这边出发的,而键 kjk_j 和值 vjv_j 全部来自编码器那边对整句中文编码出来的表示。英文每生成一个词,都会回头去中文那边按需取用对应的信息,正是交叉注意力在两种语言之间牵线搭桥。

5.注意力到底强在哪里

注意力本质上是一种按内容来找信息的办法。模型可以根据当前正在处理的内容,动态决定去参考序列里的哪些位置,而且这种参考是直接点对点的,不用像循环网络那样隔着一堆时间步慢慢传。

这就带来一个特别大的好处,哪怕是序列里隔得很远的两个词,注意力也能让它们直接发生联系,长距离依赖一下子变得好处理多了。我们回头细看一眼 3.5 节那个寝室传话游戏的窘境。循环网络里第1个词的信息想影响到第20个词,得老老实实经过中间那一大串隐藏状态的接力,每传一步都有可能被新进来的内容稀释一点、盖掉一点,传到后面早就面目全非了,这正是长距离依赖搞不定的根源。注意力机制干脆给每个位置都发了一个对讲机,第20个词想用第1个词的信息,一步直接调取就行,中间那些传递环节全省了,信息几乎不失真。

说穿了,这件事不只是模型才有的烦恼。记得有一次,小张跟我吐槽,他们组里有个项目,前面负责采集数据的同事把一条关键说明写在很早的文档里,后面接手的人一个传一个,传到最后写代码那位同事手里,意思已经完全变样了,整个项目返工了好几天。这其实就是工程版的传话游戏,注意力机制要解决的恰好就是这一类麻烦,让后面的人随时能直接翻回最早的原文核对,不必依赖中间一长串转述。

也正因为它这么管用,注意力成了下一章主角Transformer的地基,现代大模型那一身本事,相当一部分都是从这儿长出来的。

按惯例这里该收个尾,今天就先到这儿,下一章见。

练习

Q1. 注意力机制里的 query、key、value 三个角色分别是什么意思,又怎么从它们一步步算出输出?

query 是当前这一步想找什么样的信息,key 是每个候选位置身上用来被比对的标签,value 是候选位置真正能提供的内容。计算时先用 query 跟每个 key 做点积得到相似度,再过 Softmax 把这些相似度归一成一组加起来等于 1 的权重 αj\alpha_j,最后把各位置的 value 按这个权重加权和 y=jαjvjy=\sum_j\alpha_jv_j,就是输出。

Q2. 沿用本章那个手算例子:q=[1,0]q=[1,0],三个位置的 k1=[1,0],v1=[1,0]k_1=[1,0],v_1=[1,0]k2=[0,1],v2=[0,1]k_2=[0,1],v_2=[0,1]k3=[1,1],v3=[1,1]k_3=[1,1],v_3=[1,1]。请算出三个点积值,并说明第 2 个位置为什么权重最小。

qk1=1×1+0×0=1q\cdot k_1=1\times1+0\times0=1qk2=1×0+0×1=0q\cdot k_2=1\times0+0\times1=0qk3=1×1+0×1=1q\cdot k_3=1\times1+0\times1=1。三个 Softmax 权重约为 α10.422\alpha_1\approx0.422α20.155\alpha_2\approx0.155α30.422\alpha_3\approx0.422,加起来正好是 1。第 2 个位置点积为 0,是三者里最小的,说明它的 key [0,1][0,1] 跟 query [1,0][1,0] 完全不对口,所以分到的权重最小。

Q3. 缩放点积注意力里那个分母 d\sqrt{d} 是多余的装饰吗,去掉会怎样?

不是。维度 dd 越大,点积的数值就越大,大到一定程度 Softmax 的输出会几乎集中在一个位置上(其余权重被压成接近 0)。这一集中梯度就近乎为零,模型学不动。除以 d\sqrt{d} 把点积数值压回温和范围,让权重分配不那么极端,训练才稳。

Q4.(面试题) 自注意力和交叉注意力有什么区别?注意力机制相比循环网络,在处理长距离依赖上强在哪里?

自注意力里 query、key、value 都来自同一个输入序列,序列内每个位置都跟自己这批兄弟互相比对;交叉注意力则是 query 来自一个序列、key 和 value 来自另一个序列,比如机器翻译里解码器生成的每个词(query)回头去编码器那边(key、value)取用信息。相比循环网络,注意力的信息传递是直接点对点的:第 20 个词想用第 1 个词,一步直接调取就行,不用隔着一堆时间步慢慢接力,所以长距离依赖信息几乎不失真,这正是它能甩开循环网络的地方。

相关标签
深度学习注意力序列