理解 Transformer(1)

本文介绍 Transformer 从输入文本到输出预测的基本计算过程,包括 token embedding、注意力、多头结构、残差连接、encoder–decoder,以及训练和推理的区别。阅读需要矩阵乘法和梯度的基础知识。

文中以《Attention Is All You Need》的原始架构为主,并在相关位置说明它与 GPT-2 的区别。除特别说明外,省略 batch 维度,以单条序列为单位讨论。

1. 词表大小、序列长度与隐藏维度

计算中需要区分三个量:

符号 含义 决定什么
\(N_{\mathrm{vocab}}\) 词表里有多少种 token embedding 表的行数
\(S\) 当前序列有多少个 token 当前状态矩阵的行数
\(d\) 每个位置的向量维度 当前状态矩阵的列数

“100k 上下文”说的是可容纳的 token 数量上限。实际输入只有 1,000 个 token 时,当前序列长度就是 1,000。词表里的其他 token 不会因此各占一个注意力位置。

分词器有不同设计。BPE 等子词方法会从较小单位出发,把频繁出现的组合合并为 token。词表不必收录每一个完整单词:一个没见过的长词,仍可能被拆成已有片段。采用完整字节覆盖的方法可以进一步用字节表示输入;其他方法可能需要未知 token。

词表大小通常通过实验和预算取舍。词表小,文字可能切得更碎;词表大,embedding 和词表输出层更大,一些罕见 token 也更难获得充分训练。在模型训练完成后,词表和编号映射通常保持固定。

2. Token embedding

每个 token 有一个编号。编号用于查表,其大小和相邻关系不表示语义关系。设 embedding 表为:

\[E\in\mathbb{R}^{N_{\mathrm{vocab}}\times d}.\]

根据当前序列的 token 编号,从 \(E\) 中取出对应的 \(S\) 行,按序列顺序排列,得到:

\[X_{\mathrm{token}}\in\mathbb{R}^{S\times d}.\]

同一个 token 出现多次,就会重复取出同一行,放到不同的序列位置。向量的数值通过训练确定,单个分量通常没有明确的语义标签,信息可以由多个分量共同表示。

从头训练时,embedding 权重通常随机初始化,然后与整个网络一起优化。训练好的模型直接加载这张表。同一个 token 的初始向量相同,但后续会根据上下文产生不同状态。

这里需要区分两类量:参数是模型保存的可训练数值;状态是给定输入后计算出的中间结果。 \(E\) 是参数矩阵,查表得到的 \(X_{\mathrm{token}}\) 是当前输入的表示。

3. 位置编码

仅有 token embedding 还没有显式表示顺序。原论文将每个位置的编码与该位置的 token embedding 相加。设当前序列的位置编码矩阵为 \(P\),则:

\[X=\sqrt{d}\,X_{\mathrm{token}}+P,\qquad X,P\in\mathbb{R}^{S\times d}.\]

这里的 \(\sqrt{d}\) 是原论文采用的 embedding 缩放系数。相加按元素进行,结果的尺寸不变。token 编号决定从 embedding 表取哪一行;位置编号决定使用哪个位置向量。词表大小与可表示的位置数量没有必要相等。

同一个 token 出现在不同位置时,token embedding 相同,位置编码不同。相加后,网络便能利用位置差异区分不同的排列顺序。

原始 Transformer 使用固定的正弦、余弦位置编码,不同分量具有不同频率。GPT-2 使用可学习的绝对位置表。前者的位置向量由公式生成,后者的位置向量参与训练。[1, 2]

下文用 \(X\) 表示送入某个注意力模块的状态:第一层接收上述输入,后续层接收上一层处理后的结果。为具体说明尺寸,采用原论文 base 配置:\(d=512\),8 个注意力头,每头 64 维。

4. 查询、键和值

一个注意力头有三个可训练的投影矩阵:

\[W_{Q},W_{K},W_{V}\in\mathbb{R}^{512\times64}.\]

它们分别作用于同一个输入矩阵:

\[Q=XW_{Q},\qquad K=XW_{K},\qquad V=XW_{V}.\]

因此,\(Q\)、\(K\)、\(V\) 都是 \(S\times64\) 的矩阵。它们的每一行对应一个序列位置:

矩阵 含义 作用
\(Q\) 查询(query) 提供接收位置的匹配特征
\(K\) 键(key) 提供源位置的匹配特征
\(V\) 值(value) 提供后续加权求和的内容

“某个位置的 query”指 \(Q\) 中对应的那一行,是一个向量。同一个头中,所有位置都使用同一套 \(W_{Q}\)。矩阵乘法将每一行分别投影,再把结果按行排列;每个位置并没有独立的投影参数。\(K\) 和 \(V\) 同理。

\(K\) 用于计算匹配分数,\(V\) 用于后续的加权求和。分别使用两个可训练投影,使模型可以独立调整匹配所需的特征和需要传递的内容。

5. 注意力权重与加权求和

先计算所有接收位置与源位置之间的匹配分数:

\[QK^{\mathsf{T}}\in\mathbb{R}^{S\times S}.\]

上标 \(\mathsf{T}\) 表示转置。尺寸关系为:

\[(S\times64)(64\times S)=S\times S.\]

这个分数矩阵的每一行对应一个接收位置,每一列对应一个源位置。\(Q\) 的某一行与 \(K\) 的每一行分别做点积,得到 \(S\) 个分数。点积内部对向量分量求和,不同源位置的分数则分别保留。

设遮挡矩阵为 \(M\):允许读取的位置取 0,禁止读取的位置取负无穷。注意力权重为:

\[A=\operatorname{softmax}_{\mathrm{row}} \left(\frac{QK^{\mathsf{T}}}{\sqrt{64}}+M\right).\]

缩放用于控制点积分数的尺度。softmax 按行执行,将每行分数转换为非负、总和为 1 的权重;被遮挡的位置权重为零。这里假设每行至少有一个允许读取的位置。

随后用权重对值向量进行混合:

\[H=AV,\qquad (S\times S)(S\times64)=S\times64.\]

\(H\) 是这个头的输出。它的每一行,是 \(V\) 中各行按对应注意力权重形成的加权和。

输入矩阵分别投影为查询、键和值;查询与键计算匹配分数,经缩放、遮挡和 softmax 得到注意力矩阵,再与值矩阵相乘得到输出。
图 1:单个注意力头的计算过程。匹配分数决定权重,值矩阵提供被加权求和的内容。

因此,不同位置的信息通过两步影响输出:其他位置的键向量影响读取权重,其他位置的值向量参与输出的加权求和。

这也解释了同一个 token 如何在不同语境下产生不同表示。即使某个位置的初始表示相同,周围位置的键和值不同,注意力输出也可以不同。新状态继续进入后续层,但不会直接写回 embedding 表;训练时更新参数表的是反向传播得到的梯度。

将输入依赖关系显式写出:

\[H=A(X)V(X).\]

给定 \(A\) 后,加权求和对 \(V\) 是线性的;整体注意力操作则是非线性的,因为 \(A\) 也由输入计算得到。注意力权重描述这一层的信息读取,不能单独解释模型的最终判断。

6. 多头注意力

每个头都读取完整的输入矩阵,再通过各自的投影生成 \(Q\)、\(K\)、\(V\)。每个 \(512\times64\) 投影矩阵都能将全部 512 个输入分量混合到 64 个输出分量中,并不是只读取输入的一部分。

多个头的投影矩阵可以并排放成一个大矩阵,一次乘法算完,再将结果拆成各头。这与分别投影在数学上等价。

随后,各头分别计算自己的注意力权重矩阵。这样,同一个接收位置可以通过不同的头,使用不同权重读取信息。如果把所有头的通道合起来只计算一次点积和 softmax,就只得到一套位置权重。

完整输入分别进入八个注意力头,每头输出为序列长度乘64;拼接后恢复512维,再通过输出投影混合不同头的信息。
图 2:各头读取完整输入,独立计算注意力;合并后再混合不同头的输出分量。

设八个头的输出为 \(H_{1},\ldots,H_{8}\)。沿列方向拼接,再经过注意力输出投影:

\[O=\operatorname{Concat}(H_{1},\ldots,H_{8})W_{O}, \qquad W_{O}\in\mathbb{R}^{512\times512}.\]

拼接结果和 \(O\) 的尺寸都是 \(S\times512\)。\(W_{O}\) 允许不同头的输出相互混合。标准多头注意力会使用全部头的输出,各头在共同损失下训练,其功能可能互补,也可能重叠。

同一层各头的注意力计算不等待其他头的结果;下一层则可以基于上一层已经整合的信息继续处理。头数增加并行读取方式,深度增加有先后依赖的处理轮次。 若隐藏维度固定,更多头通常意味着每头更窄,并不自动带来更强能力。

7. 残差连接与 MLP

注意力输出 \(O\) 与输入 \(X\) 尺寸相同。原论文先将它们相加,再做层归一化:

\[Z=\operatorname{LayerNorm}(X+O).\]

在归一化之前,残差连接 \(X+O\) 提供了一条系数为 1 的直接通路。当更新量 \(O\) 接近零时,相加结果接近 \(X\)。这并不表示经过 LayerNorm 后的整个模块严格等于恒等变换。

用标量形式说明残差本身的梯度:

\[z=x+f(x),\qquad \frac{dz}{dx}=1+f'(x).\]

导数中的 1 对应直接通路。这有助于梯度传播,但完整网络还包含归一化等操作,不能据此保证梯度始终稳定。其他架构也可以采用缩放或门控残差。

接下来,MLP 对每一行使用同一套参数:

\[U=\operatorname{ReLU}(ZW_{1}+b_{1}),\qquad F=UW_{2}+b_{2}.\]
参数或结果 尺寸
\(Z\) \(S\times512\)
\(W_{1}\) \(512\times2048\)
\(b_{1}\) \(2048\),广播到每一行
\(U\) \(S\times2048\)
\(W_{2}\) \(2048\times512\)
\(b_{2}\) \(512\),广播到每一行
\(F\) \(S\times512\)

偏置的大小不随输入长度变化。ReLU 提供非线性,否则两个线性变换可以合成一个。MLP 在每个位置内部加工已有信息,本身不跨位置通信。然后再做一次残差与归一化:

\[X_{\mathrm{new}}=\operatorname{LayerNorm}(Z+F).\]

这就是原论文 encoder block 的主干。原论文在残差相加前等位置还使用 dropout;为简化表达,这里没有把它写进每条公式。

8. 原始 Transformer:encoder 与 decoder

前面介绍的注意力和 MLP 构成了原始 Transformer 的基本模块。完整模型由 encoder 和 decoder 两部分组成,各有六个串联的 block,同一部分内各层结构相同、参数独立。下图根据原论文 Figure 1 重绘,数据从上向下流动。[1]

本文采用的 base 配置为 512 维、8 个头。GPT-2 small 则是 768 维、12 个头和 12 个 decoder-only block;它的 MLP 激活、位置表示和归一化顺序也不同。[2]

原始 Transformer 的 encoder 和 decoder 各六层。Encoder 最终输出供每层 decoder 的交叉注意力读取。各子层均有残差连接和层归一化,decoder 最终输出再投影到词表。
图 3:原始架构的简化示意。虚线为残差;每个大框串联六次,各层参数独立。

为区分两条序列,记原文长度为 \(S_{\mathrm{src}}\),当前译文前缀长度为 \(S_{\mathrm{tgt}}\)。原文经过六层 encoder,得到:

\[C\in\mathbb{R}^{S_{\mathrm{src}}\times512}.\]

它保留原文的各个位置,并没有把整句话压成单个向量。Encoder 自注意力允许读取原文的前后位置,但不读取补齐长度用的 padding。

Decoder 的每个 block 则依次包含因果自注意力、交叉注意力和 MLP,每个子层后都有残差相加与归一化。三种注意力的区别如下:

模块 \(Q\) 的输入来源 \(K,V\) 的输入来源 可读取范围
Encoder 自注意力 当前原文状态 当前原文状态 完整原文
Decoder 自注意力 当前译文状态 当前译文状态 自己及之前的译文位置
Decoder 交叉注意力 decoder 自注意力及归一化后的状态 encoder 最终输出 \(C\) 完整原文

具体地,记 decoder 经过因果自注意力及归一化后的状态为 \(B\)。一个交叉注意力头计算:

\[Q=BW_{Q},\qquad K=CW_{K},\qquad V=CW_{V}.\]

这里的投影参数属于该交叉注意力模块,与其他注意力模块的参数独立。各矩阵尺寸为:

\[Q:S_{\mathrm{tgt}}\times64,\qquad K,V:S_{\mathrm{src}}\times64.\]

因此:

\[A:S_{\mathrm{tgt}}\times S_{\mathrm{src}}, \qquad AV:S_{\mathrm{tgt}}\times64.\]

注意力矩阵不必是方阵。输出的每一行仍对应一个译文位置,却已经读取了原文信息。多头合并和输出投影后,残差加回的是 \(B\),不是 \(C\)。

每个 decoder block 都读取同一个最终 \(C\),但使用各自的投影参数;不是第几层 decoder 对接第几层 encoder。

9. Decoder-only 模型如何读取输入

Decoder-only 模型可以把任务说明、原文和译文前缀按顺序放进同一条序列。生成译文时,原文已经位于前面,因果自注意力可以读取它。

Self-attention 的 self 指同一组序列表示,不是一个 token 只能读取自己。 同一条序列中不同位置之间的读取仍是自注意力;当查询来自 decoder,而键和值来自独立的 encoder 表示时,才是上一节介绍的交叉注意力。

因此,GPT-2 这类 decoder-only 模型不需要通过额外的交叉注意力接收文字输入。能否做好翻译还取决于训练数据与训练方式,架构提供可行的信息通道,并不自动赋予能力。

10. 推理:预测下一个 token

回到原始 encoder–decoder 模型。Decoder 最后一层的输出经词表投影和 softmax,为每个位置产生一个“下一个 token”的分布:

\[\begin{aligned} D&\in\mathbb{R}^{S_{\mathrm{tgt}}\times512},\\ W_{\mathrm{vocab}}&\in\mathbb{R}^{512\times N_{\mathrm{vocab}}},\\ \Pi&=\operatorname{softmax}_{\mathrm{row}}(DW_{\mathrm{vocab}}) \in\mathbb{R}^{S_{\mathrm{tgt}}\times N_{\mathrm{vocab}}}. \end{aligned}\]

\(D\) 是 decoder 的最终状态,\(W_{\mathrm{vocab}}\) 是词表输出投影,\(\Pi\) 是预测概率矩阵。词表输出投影与注意力内部的 \(W_{O}\) 作用不同;前者输出词表分数,后者混合多头结果。原论文还让词表输出投影与 embedding 共享权重。

推理从起始标记开始。给定当前前缀,尚未生成的下一个 token 紧跟最后一个位置,因此使用 \(\Pi\) 的最后一行。前面的状态仍参与注意力计算,为最后一个位置提供信息。

选出新 token 后加入前缀,继续预测,直到结束标记或长度限制。最直接的实现每次重算完整前缀;实际推理通常缓存各层已有位置的 \(K\)、\(V\)。因果遮挡保证旧位置不依赖后来追加的 token,因此其缓存可以复用。原文固定时,encoder 输出 \(C\) 也可复用。

选择 token 可以采用贪心、采样或束搜索。原论文的翻译实验采用束搜索,保留多个候选前缀;上述说明只追踪其中一条。

11. 训练:输入与目标错开一位

训练时有完整的原文和正确译文。设正确译文包含 \(T\) 个 token,依次记为 \(y_{1},\ldots,y_{T}\)。这里的下标只表示译文位置。

Decoder 输入与监督目标排列如下:

位置 Decoder 输入 该位置的监督目标
第 0 位 起始标记 \(y_{1}\)
第 1 位 \(y_{1}\) \(y_{2}\)
第 \(T\) 位 \(y_{T}\) 结束标记

这就是右移一位的输入。训练时 decoder 使用正确前缀,即使模型在前一个位置预测错了,也不在这次计算中用错误预测替换输入。这叫 teacher forcing。

因果遮挡让每个位置不能读取后面的答案。因此,所有位置的预测可以在一次前向计算中并行完成。得到概率矩阵后,不必先抽样生成 token,直接计算监督目标的损失。

设某个位置的正确目标为 \(y\),模型分配给它的概率为 \(p(y)\)。普通交叉熵损失为:

\[\ell=-\ln p(y).\]

正确目标的概率越高,损失越小。即使它还不是概率最高的候选,提高它的概率也会降低损失。原论文还使用了 label smoothing,对监督分布进行平滑处理。

把有效目标位置的损失汇总为 \(L\),通过链式法则反向传播,再更新参数。对任意一个可训练参数 \(w\),最简单的梯度下降更新为:

\[w_{\mathrm{new}}=w_{\mathrm{old}}-\eta\frac{\partial L}{\partial w}.\]

\(\eta\) 是学习率。一次前向、反向计算期间参数不变,算完梯度后才更新。Embedding、注意力投影、MLP、归一化中的可训练参数和输出层都可参与;\(Q\)、\(K\)、\(V\) 是中间结果,不作为独立参数更新。

训练优化的是所有样本上的共同目标。Encoder 不需要单独的目标表示 \(C\):最终预测损失的梯度会通过交叉注意力传回 encoder。各个注意力头也不需要预先指定任务。

不同样本的梯度可能相反。Batch 平均能减小单条样本带来的波动,并适合硬件并行;学习率决定每步走多远。两者是不同的控制量,不存在“越小越好”的普遍规则。原论文实际使用 Adam,而不是这里为解释方便采用的普通梯度下降。

训练时正确前缀已知,能并行计算多个位置;生成时未来 token 未知,需要逐个推进。推理时一次错误会改变后续前缀,因此,除了验证损失,还需要让模型生成完整译文来检查效果。

12. 模型容量与训练数据

参数多意味着可用容量更大,不保证已经学到了更多。数据不足时,大模型可能记住训练样本而泛化不佳;模型过小,也可能无法表达所需规律。模型大小、数据量和计算量要一起考虑,不能只数参数。

本文介绍基本计算过程。KV cache 的实现、长序列的计算成本、分布式训练和数值精度等工程问题,留待后续讨论。

参考资料与说明

  1. Vaswani et al., Attention Is All You Need, 2017。架构、尺寸和训练细节参照原论文;图 3 为重新绘制的简化示意。
  2. OpenAI, GPT-2 官方实现:model.py。用于核对 GPT-2 与原始 Transformer 的区别。
  3. 3Blue1Brown, Transformers, the tech behind LLMs。Transformer 的可视化介绍。

本文根据与 AI 助手的学习问答整理。




Enjoy Reading This Article?

Here are some more articles you might like to read next: