注意力机制与 Transformer#
学习目标与记号#
准确解释查询—键—值、缩放点积注意力与多头注意力;
根据公式和张量维度分析掩码,并识别常见实现错误;
用
Python实现或验证位置编码;在其他条件相同的情况下,只改变一个需要研究的因素进行比较,并根据实验结果分析该因素可能带来的影响。
本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;层编号写作上标 \([l]\),样本或时间编号写作下标;转置写作 \(\trans\)。除非另有说明,单个向量均为列向量;若把 \(n\) 个列向量按行组成矩阵,则统一写为 \(\bU=[\bu_1\trans;\ldots;\bu_n\trans]\)。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 注意力与 Transformer 答案。
注意力与 RNN 采用不同的上下文聚合方式;权重归一化使用 Softmax。本节结构将直接用于 BERT。
早期编码器—解码器 RNN 常把整个输入压缩到最后一个隐藏状态,形成固定长度瓶颈。注意力允许解码器在每个输出位置直接读取全部编码器状态,并根据当前查询形成不同的加权组合。输入与输出不需要等长,也不要求第 \(i\) 个源词和第 \(i\) 个目标词一一对应。
查询、键和值#
注意力可以统一概括为“用查询匹配键,再按照匹配程度汇总值”。查询表示当前位置希望寻找什么信息,键表示某个位置可以依据哪些特征与查询进行匹配,值则表示匹配后由该位置提供的实际内容。因此,查询和键用于计算“应该关注谁”,值用于确定“从被关注的位置取出什么信息”。
在自注意力中,长度为 \(T\) 的输入序列包含 \(T\) 个词元。设第 \(i\) 个词元的输入词向量为 \(\bx_i\in\mathbb{R}^{d}\),并设 \(\bW_Q,\bW_K\in\mathbb{R}^{d\times d_k}\)、\(\bW_V\in\mathbb{R}^{d\times d_v}\)。同一个词向量会分别经过三组可以通过训练确定的线性变换,得到一个查询向量、一个键向量和一个值向量:
也就是说,第 \(i\) 个词元并不是只能充当其中一种角色。在同一层自注意力中,它既可以用 \(\bq_i\) 寻找相关词元,也可以用 \(\bk_i\) 接受其他词元的匹配,还可以在受到关注时通过 \(\bv_i\) 提供信息。三个向量由不同的参数矩阵计算,因此即使它们来自同一个词向量,所表达的内容和承担的作用也不同。
为了同时表示自注意力和交叉注意力,通常用 \(T_q\) 表示查询序列包含的位置数,用 \(T_k\) 表示提供键和值的序列包含的位置数。第 \(i\) 个查询 \(\bq_i\in\mathbb{R}^{d_k}\)、第 \(j\) 个键 \(\bk_j\in\mathbb{R}^{d_k}\) 和与该键对应的值 \(\bv_j\in\mathbb{R}^{d_v}\) 均为列向量。把这些列向量转置后按行排列,可得
为什么 \(\bQ\)、\(\bK\) 和 \(\bV\) 都有以 \(T\) 表示的行数?
\(\bQ\)、\(\bK\) 和 \(\bV\) 是三个矩阵,而不是三个向量。矩阵中的一行对应一个词元产生的一个向量。在自注意力中,输入序列共有 \(T\) 个词元,每个词元都会产生一个查询向量、一个键向量和一个值向量,因此三个矩阵分别都有 \(T\) 行,即 \(T_q=T_k=T\)。这里重复出现三个 \(T\),只是分别说明三个矩阵各有 \(T\) 行,并不表示某个矩阵有 \(3T\) 行,也不表示需要把三个 \(T\) 相乘。若输入序列包含 5 个词元,则三个矩阵的行数分别是 5、5 和 5,而不是把它们合成一个有 15 行的矩阵。
从矩阵乘法也可以看出这一点。若第 \(i\) 个词元的输入表示为列向量 \(\bx_i\in\mathbb{R}^{d}\),则输入矩阵统一写为 \(\bX=[\bx_1\trans;\ldots;\bx_T\trans]\in\mathbb{R}^{T\times d}\),其每一行对应一个词元。右乘参数矩阵只会改变每一行所包含的特征数,不会改变词元的数量。因此,\(\bQ=\bX\cdot\bW_Q\in\mathbb{R}^{T\times d_k}\)、\(\bK=\bX\cdot\bW_K\in\mathbb{R}^{T\times d_k}\),而 \(\bV=\bX\cdot\bW_V\in\mathbb{R}^{T\times d_v}\)。
在交叉注意力中,查询和键值可能来自两段长度不同的序列。此时 \(\bQ\) 有 \(T_q\) 行,\(\bK\) 和 \(\bV\) 各有 \(T_k\) 行。键和值必须具有相同的行数,是因为每一个键都需要有一个与之对应的值:查询与某个键匹配后,模型才能取用该键对应的值。
缩放点积注意力为
得分矩阵的第 \((i,j)\) 个元素为 \(\bq_i\trans\cdot\bk_j/\sqrt{d_k}\),因此其维度为 \(T_q\times T_k\)。Softmax 对每个查询对应的键轴逐行计算。把各查询得到的上下文列向量按行组成 \(\bC=[\bc_1\trans;\ldots;\bc_{T_q}\trans]\in\mathbb{R}^{T_q\times d_v}\),它就是注意力机制的输出矩阵。
为什么要除以 \(\sqrt{d_k}\)?
查询向量和键向量各有 \(d_k\) 个分量,二者的点积是 \(d_k\) 个乘积之和。为了直观说明,假设各分量彼此近似独立、均值为 0 且方差为 1,则点积的方差约为 \(d_k\),其标准差约为 \(\sqrt{d_k}\)。因此,\(d_k\) 越大,点积的数值通常越容易远离 0。
如果直接把较大的正数或负数输入 Softmax,最大的得分可能获得接近 1 的权重,其余得分的权重则接近 0,导致 Softmax 过早接近“只选择一个位置”的状态,并使相应梯度变小。除以 \(\sqrt{d_k}\) 可以把点积的标准差重新调整到约为 1,使输入 Softmax 的数值范围更加稳定。这里使用的是查询向量和键向量的维度 \(d_k\);只有在 \(d_k=d\) 时,分母才等于 \(\sqrt{d}\)。
掩码 \(\bM\) 在可见位置取 0、禁止位置取一个足够小的负数,可理想化为 \(-\infty\)。这样相应 Softmax 权重为 0。若一整行全部被遮蔽,Softmax 没有有效归一化对象,程序可能产生 NaN,因此批处理时要保证掩码与有效查询一致。
图 28 编码器—解码器注意力示意图#
在 RNN 编码器—解码器中,解码器状态可产生查询,编码器各时间步状态产生键和值。第 \(i\) 个解码位置的上下文向量为
相似度 \(e_{ij}\) 可以是点积、双线性形式 \(\bq_i\trans\cdot\bW\cdot\bk_j\),或加性注意力 \(\bv\trans\cdot\tanh(\bW_q\cdot\bq_i+\bW_k\cdot\bk_j)\)。注意力权重描述模型当前的计算分配,不应自动等同于因果解释或人类可读的“重要性”。
自注意力与位置表示#
设第 \(i\) 个词元的输入表示 \(\bx_i\in\mathbb{R}^{d}\) 为列向量,并将其转置后按行堆叠为 \(\bX=[\bx_1\trans;\ldots;\bx_T\trans]\in\mathbb{R}^{T\times d}\)。令 \(\bW_Q,\bW_K\in\mathbb{R}^{d\times d_k}\)、\(\bW_V\in\mathbb{R}^{d\times d_v}\),则单个位置的投影满足 \(\bq_i=\bW_Q\trans\cdot\bx_i\)、\(\bk_i=\bW_K\trans\cdot\bx_i\) 和 \(\bv_i=\bW_V\trans\cdot\bx_i\)。把所有位置一起计算,可得
不加入位置表示时,自注意力对输入位置的共同置换是等变的,本身不知道“第几个词”。一种固定正弦位置编码使用列向量 \(\bp_p\in\mathbb{R}^{d}\) 表示位置 \(p\),其分量为
其中 \(p\) 是位置,\(i\) 是频率下标。位置编码矩阵统一写为 \(\bP=[\bp_1\trans;\ldots;\bp_T\trans]\in\mathbb{R}^{T\times d}\)。把每个位置列向量与同一位置的词元嵌入相加,再将结果的转置按行组成序列矩阵,模型便能利用位置信息。位置表示也可以学习,或通过旋转、相对偏置等方式直接作用于注意力得分;关键是明确所用方法的长度外推性质和掩码约定。
多头注意力#
多头注意力让不同子空间分别计算注意力。设查询侧输入、键侧输入和值侧输入的列向量分别为 \(\bx_i^Q\)、\(\bx_j^K\) 和 \(\bx_j^V\),则相应序列矩阵统一写为 \(\bX_Q=[(\bx_1^Q)\trans;\ldots;(\bx_{T_q}^Q)\trans]\)、\(\bX_K=[(\bx_1^K)\trans;\ldots;(\bx_{T_k}^K)\trans]\) 和 \(\bX_V=[(\bx_1^V)\trans;\ldots;(\bx_{T_k}^V)\trans]\)。第 \(r\) 个头为
最终输出为
常见设置令每头 \(d_k=d_v=d/H\),使拼接后维度仍为 \(d\),但这不是定义所强制的。多个头扩大了可学习的关系类型,却不保证每个头都会获得清晰、彼此不同的语义。
掩码的职责#
键填充掩码:阻止有效查询读取
[PAD]位置。还应在损失中屏蔽填充目标;仅遮蔽键并不会自动删除填充查询的输出。因果掩码:自回归解码位置 \(t\) 只能读取不晚于 \(t\) 的键。训练时可以并行计算所有位置,但依赖图仍保持因果;生成时因为下一个词元尚未出现,通常仍需逐步解码和缓存键值。
任务掩码:根据结构限制可见连接,例如局部窗口或特定模态之间的连接。BERT 的
[MASK]是输入词元替换策略,不应与注意力矩阵中的可见性掩码混为一谈。
图 29 因果注意力掩码示意图#
Layer Normalization#
现以一个批量为例说明 Layer Normalization 的计算。本小节用 \(m\) 表示一个批量中的样本数。简单起见,设每个样本均包含 \(T\) 个词元,每个词元用 \(d\) 个特征表示。把整个批量记为 \(\bX\in\mathbb{R}^{m\times T\times d}\),其中 \(x_{i,t,j}\) 表示第 \(i\) 个样本、第 \(t\) 个词元的第 \(j\) 个特征,且 \(i=1,\ldots,m\)、\(t=1,\ldots,T\)、\(j=1,\ldots,d\)。
Layer Normalization 对每个样本的每个词元分别进行计算。对于第 \(i\) 个样本的第 \(t\) 个词元,只使用该词元自身的 \(d\) 个特征计算均值和方差:
其中,\(\mu_{i,t}\) 和 \(v_{i,t}\) 分别是第 \(i\) 个样本中第 \(t\) 个词元在特征维度上的均值和方差。这里使用分母 \(d\),而不是 \(d-1\),因为目的在于归一化当前词元的特征,而不是根据样本估计总体方差。
第 \(j\) 个特征的归一化输出为
其中,\(\epsilon>0\) 是防止分母为 0 的很小正数,\(\bgamma=(\gamma_1,\ldots,\gamma_d)\trans\in\mathbb{R}^{d}\) 和 \(\bbeta=(\beta_1,\ldots,\beta_d)\trans\in\mathbb{R}^{d}\) 分别是可以通过训练确定的缩放参数和平移参数。同一组 \(\bgamma\) 和 \(\bbeta\) 用于批量中的全部 \(m\) 个样本及每个样本的全部 \(T\) 个位置。输出张量 \(\bY\) 与输入张量的维度相同,即 \(\bY\in\mathbb{R}^{m\times T\times d}\)。
在应用 \(\bgamma\) 和 \(\bbeta\) 之前,归一化后 \(d\) 个特征的均值为 0,平方均值为 \(v_{i,t}/(v_{i,t}+\epsilon)\);当 \(\epsilon\) 相对于 \(v_{i,t}\) 很小时,该平方均值接近 1。应用缩放和平移之后,最终输出的均值不必为 0,方差也不必为 1,因为模型可以通过训练 \(\bgamma\) 和 \(\bbeta\) 调整各个特征的尺度和位置。
批量大小 \(m\) 是否参与均值和方差的计算?
不参与。一个批量中共有 \(mT\) 个词元位置,Layer Normalization 会分别为这些位置计算 \(mT\) 组均值和方差,每一组统计量都只来自相应词元的 \(d\) 个特征。批量大小 \(m\) 只表示一次并行处理多少个样本,不会出现在 \(\mu_{i,t}\) 和 \(v_{i,t}\) 的求和范围或分母中。因此,即使改变同批的其他样本,第 \(i\) 个样本中第 \(t\) 个词元的归一化结果也不会随之改变。
Layer Normalization 在训练和预测时都按照上述相同公式,使用当前词元的 \(d\) 个特征进行计算,不需要保存整个训练过程中的均值和方差。相比之下,Batch Normalization 在训练时会使用当前批量的统计量,预测时通常使用训练期间逐步记录的统计量。对于包含填充位置的变长序列,Layer Normalization 仍会对填充位置进行计算,但该位置不会参与其他有效位置的均值和方差;模型仍需使用注意力掩码和损失掩码,避免把填充位置当作真实词元使用。
Transformer 层#
一个 Transformer 编码器层包含多头自注意力和逐位置前馈网络。若每个词元表示 \(\bx_t\in\mathbb{R}^{d}\) 都是列向量,并令 \(\bX=[\bx_1\trans;\ldots;\bx_T\trans]\in\mathbb{R}^{T\times d}\) 将它们的转置按行存放,则
其中 \(\bone\in\mathbb{R}^{T}\) 是元素全为 1 的列向量,\(\bb_1\) 和 \(\bb_2\) 也都是列向量;因此 \(\bone\cdot\bb_1\trans\) 和 \(\bone\cdot\bb_2\trans\) 会把相应偏置的转置复制到每一个词元位置。FFN 对每个位置独立使用同一参数。每个子层还配有残差连接、Dropout 和 LayerNorm。Vaswani 等(2017) 的原始论文采用 post-norm,现代模型也常采用更利于深层优化的 pre-norm,例如
阅读结构图或代码时必须确认归一化在残差分支之前还是之后。
编码器使用非因果自注意力;解码器包含因果自注意力,以及用解码器状态作查询、编码器输出作键和值的交叉注意力。最后通过词表投影
其中,隐藏状态 \(\bh_t\in\mathbb{R}^{d}\) 与输出 \(\bz_t\in\mathbb{R}^{|\mathcal{V}|}\) 均为列向量,\(\bW_{\mathrm{vocab}}\in\mathbb{R}^{|\mathcal{V}|\times d}\),\(\bb_{\mathrm{vocab}}\in\mathbb{R}^{|\mathcal{V}|}\),而 \(|\mathcal{V}|\) 表示词表大小。这个线性投影不是 Embedding 的逆运算;有些模型会让投影权重与输入嵌入共享,但共享不是 Transformer 定义的必要条件。
计算代价与适用范围#
标准全局自注意力需要显式或隐式处理 \(T\times T\) 得分,时间和注意力内存通常随序列长度呈二次增长。它缩短了任意两位置之间的计算路径,并能在训练时并行处理所有查询,但长序列成本可能高于 RNN。局部、稀疏或线性化注意力可降低成本,同时也改变可见范围或近似性质。
Shiny 交互演示:自注意力与 Layer Normalization
交互页面按照正文公式计算查询、键、值、缩放点积得分、逐行 Softmax 权重和注意力输出,并可切换因果掩码。页面还逐位置展示 Layer Normalization 的均值、方差、标准化结果以及尺度和平移。随机矩阵只用于核对公式和维度,注意力权重也不应自动解释为因果关系。
核心推导与实现核验#
核心关系
推导路径。 查询与键点积形成相似度,按键轴做 Softmax 得到每个查询的权重,再对值加权求和;除以 \(\sqrt{d_k}\) 控制随机点积方差。
关键条件
若查询和键各分量独立、零均值、单位方差,点积是 \(d_k\) 项之和,方差为 \(d_k\);缩放后方差约为 1,可减少 Softmax 过度饱和。
数据规模
若批次大小为 \(m\)、注意力有 \(H\) 个头、查询位置有 \(T_q\) 个、键位置有 \(T_k\) 个,则注意力分数表的大小为 \(m\times H\times T_q\times T_k\)。掩码必须能覆盖这张表,并明确每个查询位置允许查看哪些键位置。
常见误区
把 Softmax 放在查询轴会让权重解释错误;在 Softmax 后再置零而不重新归一化也会使行和小于 1。
动手检查
检查注意力权重沿键轴和为 1、被屏蔽位置严格为 0;将所有值向量设为同一常量时,每个查询输出应等于该常量。
数值稳定性与规模
分数除以 sqrt(d_k),沿键轴减最大值;无效位置在 Softmax 前填充负无穷,并确保每个查询至少有一个有效键。
本节小结#
注意力权重必须沿键位置归一化。
缩放项控制点积方差与 Softmax 饱和。
多头拆分、掩码和位置表示都需严格检查轴。
综合练习#
程序题应固定随机种子、写出维度断言并报告运行环境。全部参考答案见 注意力与 Transformer 答案。
注意力数值计算。 令 \(\bQ=\bK=\bI_2\)、\(\bV=\begin{bmatrix}2&0\\0&4\end{bmatrix}\),其中三个矩阵的每一行分别表示相应查询、键和值列向量的转置,且无掩码。依次计算缩放后的得分矩阵、逐行 Softmax 权重和最终输出,保留至少四位小数。
权重性质证明。 证明每个至少含一个有效键的查询,其注意力权重非负且沿键轴和为 1;进而证明当所有有效值向量均等于同一向量 \(\bv\) 时,该查询的输出仍为 \(\bv\)。说明全被遮蔽的查询为何不满足推导前提。
缩放因子推导。 假设查询和键的各分量相互独立、均值为 0、方差为 1。推导点积 \(\bq\trans\cdot\bk\) 的均值和方差,再证明除以 \(\sqrt{d_k}\) 后方差为 1。指出真实网络中哪些条件可能不完全成立。
多头维度与内存计算。 对 \(m=8\)、\(T=128\)、\(d=256\)、\(H=8\) 的自注意力,写出每头 \(d_k\),以及 \(\bQ\)、\(\bK\)、\(\bV\)、注意力得分和拼接输出的批量维度。计算所有头的得分元素数及单个 float32 得分张量的内存,并写出标准注意力关于 \(T\) 的主要时间和内存数量级。
带掩码注意力实现。 实现支持批次、多头、键填充掩码和因果掩码的稳定缩放点积注意力。Softmax 前沿键轴减去有效得分最大值,被屏蔽权重必须严格为 0;程序应拒绝全被遮蔽的有效查询,并返回输出和注意力权重。
公式与库实现核验。 按照 手算结果与可信实现对照 的原则,用小矩阵逐项核对第 5 题的得分、权重和输出,并与
torch.nn.functional.scaled_dot_product_attention或等价可信实现比较。测试无掩码、仅 padding、仅 causal、两种掩码同时存在和极大得分,检查权重行和、屏蔽位置、输出维度及梯度。循环、批量与位置性质。 写出逐查询循环版和批量矩阵版单头注意力,验证二者结果一致并比较运行时间;再验证无位置表示的自注意力对共同位置置换满足等变性,而加入固定位置编码后一般不再满足该性质。
RNN 与 Transformer。 在同一序列分类任务上比较单层 GRU 与小型 Transformer 编码器。固定训练/验证/测试划分、随机种子集合、输入词表、近似参数量、训练更新次数和每次更新的样本数;报告预测质量、训练时间、固定批量预测时间、峰值内存、参数量和网络层调用次数,并按序列长度分组报告结果。
注意力头数比较。 在固定模型维度 \(d\) 下比较 \(H\in\{1,2,4,8\}\)。固定数据划分、随机种子集合、Transformer 层数、前馈维度、训练更新次数和评价程序;报告预测质量、训练时间、固定批量预测时间、峰值内存、参数量和注意力模块调用次数,检查 \(d\) 是否能被 \(H\) 整除。
全局与局部注意力。 在长序列任务上比较标准全局注意力和多个滑动窗口宽度的局部注意力。各组使用相同的数据划分、随机种子、模型宽度与层数、训练更新次数、总处理词元数和评价程序,只改变注意力窗口宽度;报告预测质量、训练时间、固定批量预测时间、峰值内存、注意力得分元素数和网络函数计算次数。检查完成任务所需的信息是否超出窗口范围,据此说明窗口宽度为什么可能影响预测结果和计算开销。