我第一次手写Transformer编码器的时候,最大的感受不是“这个模型有多复杂”,而是“组件明明就这几块,怎么一拼起来就到处报维度错误”。调包调习惯了,真到要自己从零实现一遍,才发现很多细节不是看论文就能看出来的——尤其是多头注意力里那些reshape、transpose、mask广播,每一步踩错都够你折腾半小时。这篇就把Transformer的编码器部分拆开揉碎讲清楚,从输入表示、位置编码、多头自注意力到前馈网络和层归一化,全部用PyTorch手写实现,不用现成的torch.nn.Transformer,这样里面的机制才能真正长在自己脑子里。如果你已经会PyTorch的基本操作,又不想一直停留在调包阶段,这篇应该能帮你把编码器的每一个齿轮都转明白。
1. 从整体架构定位编码器
1.1 Transformer不是“一个模型”,而是一套框架
Transformer是2017年《Attention Is All You Need》里提出的序列到序列模型,最初用在机器翻译上。它的整体结构可以理解成一个“编码器—解码器”的管道:编码器负责把源语言的句子读进去,转成一组携带上下文信息的特征向量;解码器再根据这些特征向量,一个一个地生成目标语言的词。
这里要特别强调一句:如果你在搜索引擎里搜“编码器”,大概率还会搜到工业上测旋转角度的光电编码器、磁编码器,那是完全不同的东西。这篇文章里说的编码器,是Transformer左边的Encoder模块。
Encoder模块做的事情,用一句话概括就是:输入一个词序列,输出一个等长的、每个位置都融合了全文信息的向量序列。比如输入“小明打篮球,他很开心”,编码器最终输出的每个向量,都不再是某个词的孤立表示,而是“看过整句话之后的这个词的表示”。也正是这一点,让后续的BERT这类预训练模型有了立足之地——它们本质上都是在用Transformer的编码器去学习文本的上下文表示。
1.2 编码器的内部结构:一层搞定所有核心机制
编码器不是一个巨大的单一网络,而是N个结构完全相同的层堆叠在一起。原论文里N=6,每一层内部包含两个子层:
- 第一个子层:多头自注意力机制(Multi-Head Self-Attention)
- 第二个子层:位置逐前馈网络(Position-Wise Feed-Forward Network)
每个子层外面都套了一层残差连接和LayerNorm,也就是经典的“Add & Norm”结构。
| 组件 | 输入 | 输出 | 作用 |
|---|---|---|---|
| 多头自注意力 | 上一层输出[B, L, D] | 上下文融合后的[B, L, D] | 让每个token与其他token交换信息 |
| 前馈网络 | 注意力输出[B, L, D] | 非线性变换后的[B, L, D] | 对每个token做独立的特征变换 |
| 残差连接 + LayerNorm | 子层输入与输出之和 | 归一化后的[B, L, D] | 缓解深层网络梯度消失,稳定训练 |
这里的L是序列长度,D是特征维度。整个编码器不管有多少层,每一层的输入输出形状都保持一致,这也是它能方便堆叠的原因。
1.3 为什么建议先手写编码器,再碰解码器
很多人一上来就想直接实现整个Transformer,结果被解码器的Masked Self-Attention和Cross-Attention搞得晕头转向。我的建议是:先把编码器写稳了再说。原因很简单,Transformer里最核心的计算逻辑——多头自注意力、残差、LayerNorm、前馈网络——在编码器里全都有,而解码器只是在编码器的基础上加了两个东西:
- 在自注意力里加因果Mask,让当前位置只能看到它之前的token,保证生成时不会“偷看未来”。
- 中间插入一个Cross-Attention子层,让解码器可以去查询编码器的输出。
先把编码器吃透,后面理解解码器就只是“加Mask”和“多一个注意力模块”的事,不会有本质性的新概念。这也是这篇“上篇”只讲编码器的原因。
2. 输入表示:Embedding与位置编码
2.1 Token Embedding:把词变成向量
Transformer不能直接处理离散的词,第一步要做的是查表,把每个词的索引映射成一个稠密向量。在PyTorch里就是nn.Embedding,输入形状是[batch, seq_len],输出形状是[batch, seq_len, d_model]。
一个很容易被忽略的细节是:原论文在Embedding之后对向量乘了一个sqrt(d_model)的缩放因子。为什么要这么做?因为后面的位置编码是直接加到Embedding向量上的,而位置编码的数值范围是固定的[-1, 1],如果不放大Embedding,位置编码的信息占比会偏高,可能干扰词本身的语义表示。这个细节在很多简化实现里被省略了,但严谨起见建议保留。
import math import torch import torch.nn as nn import torch.nn.functional as F class TokenEmbedding(nn.Module): def __init__(self, vocab_size: int, d_model: int): super().__init__() self.embed = nn.Embedding(vocab_size, d_model) self.d_model = d_model def forward(self, x): # x: [batch, seq_len] return self.embed(x) * math.sqrt(self.d_model)2.2 位置编码:为什么Transformer必须“有位置感”
RNN天生是按顺序处理序列的,词与词的先后关系天然被编码在时间步里;CNN也能通过卷积核的感受野在一定程度上感知局部位置。但Transformer的自注意力机制是一个“全连接”结构——它计算两个token的相关性时,完全不关心它们在序列里的先后顺序。
打个比方,让模型处理“小明打篮球”和“篮球打小明”,如果只看词的集合,这两个句子的词完全一样,但意思完全相反。如果没有位置信息,自注意力会把这两句话表示成一模一样的向量,这肯定是不可接受的。
所以,Transformer在把输入送进编码器之前,必须显式地给每个位置加上一个“位置标记”。位置编码就是干这个的。
2.3 原论文的正余弦位置编码实现
原论文用的是固定公式的正余弦位置编码,不参与训练:
- 对位置pos、维度索引2i:
PE(pos, 2i) = sin(pos / 10000^(2i / d_model)) - 对位置pos、维度索引2i+1:
PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))
乍看这个公式有点吓人,但拆开看其实很清晰:它给每个位置生成一个长度为d_model的向量,不同维度对应不同的正弦/余弦波频率。低维的波频率高,高维的波频率低。
为什么用正余弦而不是直接学一个位置向量?原论文给的理由是:正余弦形式可以让模型容易地通过线性变换来表达相对位置关系,而且当测试时遇到比训练时更长的序列,公式可以直接外推,不需要额外训练。
这里还有一个工程上的细节:pe需要通过register_buffer注册,这样它会被自动移入GPU,但不会参与梯度更新。
class PositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int = 5000, dropout: float = 0.1): super().__init__() self.dropout = nn.Dropout(p=dropout) pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) # 用指数形式计算 1 / 10000^(2i / d_model) div_term = torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) pe = pe.unsqueeze(0) # [1, max_len, d_model] self.register_buffer('pe', pe) def forward(self, x): # x: [batch, seq_len, d_model] x = x + self.pe[:, :x.size(1)] return self.dropout(x)代码里的div_term和公式是等价的:exp(-log(10000) * i / d_model)就等于1 / 10000^(i / d_model),只是数值上更稳定,不容易出现浮点数上溢或下溢。
有一点要特别注意:位置编码跟输入序列长度无关,pe是预先生成好长度为max_len的完整矩阵,计算时只截取当前序列长度那一段self.pe[:, :x.size(1)],这样变长序列也能正确处理。
3. 多头自注意力机制与实现
3.1 自注意力:让每个token去看其他token
自注意力是Transformer的“灵魂”。它的目标很朴素:对于序列中的每一个token,计算它和其他所有token的相关性,然后用相关性去加权整合其他token的信息。
要理解这个计算过程,可以借助一个“人”的比喻:假设每个token是一个人,它有三个角色——
- Q(Query):代表“我在找什么”。相当于你在人群中喊“谁负责这个知识点?”。
- K(Key):代表“我有什么特征”。相当于每个人身上挂的标签“我负责这个知识点”。
- V(Value):代表“如果匹配上了,我提供什么信息”。相当于这个人真正能告诉你的内容。
计算注意力时,每个token先用Q去和所有token的K做点积。点积越大,说明两个向量越相似,也就是这个人越匹配。然后对点积结果做softmax,得到一组权重,最后用权重对所有V做加权求和。加权求和后的向量,就是这个token看完整个序列之后得到的“上下文表示”。
3.2 为什么要除以根号d_k
注意力分数公式是:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V
这里的d_k是每个注意力头的维度。很多人会忽略除以sqrt(d_k)这一步,觉得“不除也能跑”,确实能跑,但效果和稳定性会差很多。
原因要从概率统计角度看:如果Q和K里的元素是均值为0、方差为1的随机变量,那它们点积结果的方差约等于d_k,标准差是sqrt(d_k)。当d_k比较大的时候,点积的数值会很大,softmax的输入分布会变得很“尖锐”——某个值特别大,其他值特别小,softmax输出会接近one-hot,梯度趋近于0,模型几乎学不动。
除以sqrt(d_k)之后,点积的方差被拉回1附近,softmax的输入分布保持在合适的区间,梯度能顺畅回传。这一步不是拍脑袋加的,是让自注意力能稳定训练的关键操作。
3.3 多头:不是只看一个“最相关”,而是看多种相关
单头注意力的问题是:它只能学一种“相关性”。但真实语言里的相关性是多种多样的——有的头应该关注代词和先行词的指代关系,有的头应该关注“形容词修饰哪个名词”,有的头应该关注句法结构。
多头注意力的思路是:把特征维度切成h份,每一份独立做一次注意力计算,这样每个头可以在不同的表示子空间里学习不同的注意力模式。最后把h个头的输出拼起来,再过一层线性投影,融合所有头的信息。原论文用的是8个头,d_model=512,每个头维度是64。
3.4 MultiHeadAttention的完整PyTorch实现
实现多头注意力时,最容易搞混的就是维度变换。我习惯把整个流程分成三步:
- 对输入做Q/K/V线性投影,得到
[batch, seq_len, d_model]。 - 把
d_model拆成num_heads * head_dim,然后转置成[batch, num_heads, seq_len, head_dim]。 - 在
seq_len维度上做矩阵乘法,算注意力分数和加权结果,最后再转置回原来的形状。
代码如下:
class MultiHeadAttention(nn.Module): def __init__(self, d_model: int, num_heads: int, dropout: float = 0.1): super().__init__() assert d_model % num_heads == 0, "d_model必须能被num_heads整除" self.d_model = d_model self.num_heads = num_heads self.head_dim = d_model // num_heads self.w_q = nn.Linear(d_model, d_model) self.w_k = nn.Linear(d_model, d_model) self.w_v = nn.Linear(d_model, d_model) self.w_o = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(p=dropout) def forward(self, x, mask=None): batch, seq_len, _ = x.shape # 1. 线性投影并拆多头 q = self.w_q(x) # [B, L, D] k = self.w_k(x) v = self.w_v(x) # 2. 将最后一维拆成 [num_heads, head_dim],并转置为 [B, h, L, d_k] q = q.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k = k.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) v = v.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) # 3. 缩放点积注意力 scale = 1.0 / math.sqrt(self.head_dim) scores = torch.matmul(q, k.transpose(-2, -1)) * scale # [B, h, L, L] if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) attn = torch.softmax(scores, dim=-1) attn = self.dropout(attn) out = torch.matmul(attn, v) # [B, h, L, d_k] # 4. 合并多头,线性输出 out = out.transpose(1, 2).contiguous().view(batch, seq_len, self.d_model) return self.w_o(out)这里有几个细节值得单独拎出来说。
第一,缩放因子用的是head_dim而不是d_model。因为实际参与点积的是每个头的head_dim维向量,不是整个d_model,除错了尺度就错了。
第二,view之前必须先保证张量在内存中是连续的。transpose操作不会复制数据,只是改变了“视图的解读方式”,这时候直接view往往会报错。代码里在transpose(1, 2)之后、view(batch, seq_len, self.d_model)之前加了.contiguous(),目的就是让张量在内存中真正排成我们希望的样子。
第三,mask的值用的是-1e9而不是0。因为softmax里如果把一项填成0,经过指数运算依然会是1,起不到屏蔽作用;填一个很大的负数,softmax之后这一项才会趋近于0。有些实现直接填float('-inf'),也可以,但-1e9在数值上更保守,不容易在某些极端情况下产生NaN。
对于mask的维度,我这里默认传入的是已经能广播到[batch, num_heads, seq_len, seq_len]的形状。实际使用时,如果你有一个[batch, seq_len]的padding mask,需要先unsqueeze(1).unsqueeze(1)变成[batch, 1, 1, seq_len],这样广播机制会自动扩展到所有头。这一点在后面的完整代码里会看到。
4. 前馈网络与残差归一化的设计考量
4.1 位置前馈网络:每个token独立做一次非线性变换
自注意力机制本质上是一个“加权求和”操作,是线性的。如果整个编码器只有自注意力,那不管堆多少层,从数学上看都还是线性变换,模型的表达能力会非常有限。所以每一个编码器层里都要配一个非线性变换模块——前馈网络。
“位置前馈网络”这个名字听起来高大上,实际就是两个线性层加一个激活函数,中间先把维度从d_model放大到d_ff,再压缩回d_model。原论文里d_ff = 2048,是d_model = 512的4倍。
为什么要先放大再缩小,而且放大到4倍?这是一个经验设计。放大维度给了模型一个临时的“工作记忆空间”,可以在这个高维空间里做更复杂的特征变换,再压缩回原来的维度;4倍是一个在计算量和表达能力之间比较平衡的选择,后续很多大模型也沿用这个比例。
class FeedForward(nn.Module): def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1): super().__init__() self.linear1 = nn.Linear(d_model, d_ff) self.linear2 = nn.Linear(d_ff, d_model) self.dropout = nn.Dropout(p=dropout) def forward(self, x): return self.linear2(self.dropout(F.relu(self.linear1(x))))激活函数的选择上,原论文用的是ReLU,现在很多实现(尤其是GPT系列)会换成GELU,效果普遍略好,训练也更稳定。在小规模任务上两者差别不大,ReLU实现更简单,GELU收敛通常更平稳一些。你可以把激活函数部分抽成一个参数,方便后续替换。
4.2 残差连接:给梯度开一条“高速公路”
Transformer的编码器要堆6层、甚至更多层。如果没有残差连接,深层网络的梯度需要通过每一层的矩阵乘法逐层回传,很容易越传越小,最终梯度消失。残差连接做的事情很简单:把子层的输入直接加到子层的输出上。
x = x + sublayer(x)
这样做的好处是,梯度在反向传播时可以直接通过“加法捷径”流回前面,不需要穿透整个子层的变换矩阵。每一层在这个基础上只需要学习“输入和输出之间的差值”,也就是残差,学习难度大大降低。从信息流的角度看,深层的表示依然保留了一部分原始输入的信息,不会因为逐层变换而丢失太多底层特征。
4.3 LayerNorm与BatchNorm:为什么Transformer选LayerNorm
很多人刚接触Transformer时都会问:为什么不用BatchNorm?
BatchNorm是在整个batch的维度上做归一化,统计每个特征的均值和方差,这会导致两个问题:
- 对batch size敏感。batch太小,统计量不稳定,训练容易波动。
- 对序列长度敏感。NLP里同一个batch的句子长度往往不一致,padding之后,不同位置的统计量会受到padding影响,处理起来很别扭。
LayerNorm不一样,它对每个样本、每个token单独做归一化,只统计当前token特征维度的均值和方差,完全不依赖batch,也不依赖其他token。因此不管batch多大、句子多长,LayerNorm的统计量都是稳定的。这一点对变长序列和在线推理特别友好。
| 对比项 | BatchNorm | LayerNorm |
|---|---|---|
| 归一化维度 | 每个特征维度跨batch归一化 | 每个样本/每个token内特征维度归一化 |
| 对batch size的依赖 | 依赖,batch小不稳定 | 不依赖 |
| 对序列长度的处理 | 容易受padding影响 | 不受序列长度影响 |
| 训练稳定性 | 在NLP任务中容易波动 | 稳定,适合Transformer |
在代码里,LayerNorm的实现很简单:nn.LayerNorm(d_model),它会自动学习仿射变换参数(缩放和平移),不需要我们手动管理。
还有一个常被讨论的细节:原论文用的是Post-LN,也就是“残差相加之后再归一化”——这也是我下面代码里采用的方式。但实际操作中,Pre-LN(先归一化,再进子层,最后残差相加)在深层模型里训练更稳定,很多现代框架默认用Pre-LN。对小规模任务来说,两者都能收敛,不用过分纠结;我自己的经验是,如果层数超过12层,Pre-LN会让你少调很多学习率相关的参数。
5. 完整编码器forward流程与PyTorch实现
5.1 EncoderLayer:把注意力、前馈、归一化组装起来
有了上面的基础模块,编码器层就很简单了。它本质上就是把多头自注意力、前馈网络、残差连接、LayerNorm按顺序组合在一起。每个子层都遵循同一个套路:
- 先做子层计算
- 加上残差输入
- LayerNorm归一化
代码如下:
class EncoderLayer(nn.Module): def __init__(self, d_model: int, num_heads: int, d_ff: int, dropout: float = 0.1): super().__init__() self.attn = MultiHeadAttention(d_model, num_heads, dropout) self.ffn = FeedForward(d_model, d_ff, dropout) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(p=dropout) def forward(self, x, mask=None): # 第一个子层:多头自注意力 + 残差 + LayerNorm attn_out = self.attn(x, mask) x = self.norm1(x + self.dropout(attn_out)) # 第二个子层:前馈网络 + 残差 + LayerNorm ffn_out = self.ffn(x) x = self.norm2(x + self.dropout(ffn_out)) return x有一个地方要解释一下:我在残差相加之前对子层输出过了Dropout。这是原论文和PyTorch官方Transformer实现里的常见做法,目的是在残差相加前随机丢弃一部分子层输出,起到正则化作用,防止模型过拟合。如果你用Pre-LN结构,Dropout一般放在子层输出之后、残差相加之前,位置也差不多。
5.2 Encoder:多层堆叠与完整forward流程
一个完整的编码器由三部分组成:TokenEmbedding、PositionalEncoding、N个EncoderLayer。forward流程是:
- 输入token索引 -> Embedding
- 加上位置编码
- 依次通过N个EncoderLayer
- 输出上下文特征表示
class Encoder(nn.Module): def __init__( self, vocab_size: int, d_model: int = 512, num_heads: int = 8, d_ff: int = 2048, num_layers: int = 6, max_len: int = 5000, dropout: float = 0.1, ): super().__init__() self.embed = TokenEmbedding(vocab_size, d_model) self.pos_enc = PositionalEncoding(d_model, max_len, dropout) self.layers = nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) def forward(self, tokens, mask=None): x = self.embed(tokens) x = self.pos_enc(x) for layer in self.layers: x = layer(x, mask) return x用随机数据测一下:
encoder = Encoder(vocab_size=1000, d_model=512, num_heads=8, d_ff=2048, num_layers=6) x = torch.randint(0, 1000, (2, 20)) out = encoder(x) print(out.shape) # torch.Size([2, 20, 512])输出形状和输入序列长度保持一致,特征维度还是512。这说明编码器做的是“序列到序列”的变换,只不过这个输出序列里的每个向量,都已经包含了全序列的上下文信息。
5.3 简化实现与论文原版之间,有哪些可以省略的细节
网上很多Transformer教程会去掉Embedding,直接用一个随机向量作为输入,甚至连位置编码都省略。这种“最小实现”优点是简单,适合理解核心机制,但严格照着做,你会发现它跑小任务可能没问题,但放到真实数据集上性能会差不少。
我个人的建议是,认真复现时至少要保留这几样:
- Embedding后的
sqrt(d_model)缩放。 - 位置编码(不管是正余弦还是可学习的,必须有一个)。
- 每个子层之后的Dropout和LayerNorm。
至于max_len长度、激活函数选ReLU还是GELU、是否用Pre-LN替代Post-LN,这些属于可以按任务灵活调整的“超参数”。初学者不需要一上来就追求完全复现论文,先把这个结构跑通,再逐步替换细节,效果会更扎实。
6. 训练中的维度调试与常见错误排查
6.1 我反复踩过的3个维度坑
手写编码器时,最让人崩溃的往往不是原理,而是维度不对。以下三个坑我帮大家提前踩过了。
第一个坑:多头拆分后忘记转置。把d_model拆成num_heads * head_dim之后,张量形状是[B, L, h, d_k],但要做矩阵乘法必须把形状变成[B, h, L, d_k],也就是要在第1和第2维之间transpose(1, 2)。漏掉这一步,你在torch.matmul(q, k.transpose(-2, -1))时会发现矩阵乘法怎么都说不通,或者结果形状完全不对。
第二个坑:transpose之后直接view。transpose只是改变了张量的视图,内存布局没变,此时直接view会报RuntimeError: view size is not compatible with input tensor's size and stride。解决办法是.contiguous().view(...),或者干脆用.reshape(...)——reshape在需要时会自动复制数据,省心一些。我后来写代码几乎统一用reshape,避免在调试时被这种“隐形问题”绊住。
第三个坑:mask维度没有广播到位。很多场景下你手里的mask是[batch, seq_len],但attention scores的形状是[batch, num_heads, seq_len, seq_len]。如果不做unsqueeze(1).unsqueeze(1),广播就会失败或者得到错误结果。记好这个固定套路:mask = mask.unsqueeze(1).unsqueeze(1),让mask变成[batch, 1, 1, seq_len],它就能自动广播到所有头和所有token。
6.2 怎么验证编码器实现是对的
代码写完之后,怎么确认它真的对的?我建议从三个层次验证。
第一层:跑前向。输出形状必须和输入序列长度一致,特征维度保持d_model。这一步能挡住绝大多数维度错误。
第二层:手工验证注意力分数。拿一个极简例子,比如序列长度2、一个头、d_model=2,手动算一遍QK^T除以sqrt(d_k)之后的数值,再用代码输出对比一下。这个方法虽然原始,但能非常有效地确认softmax之前的逻辑有没有写错。
第三层:反向传播不报错,且loss能下降。你用随机输入跑一次loss.backward(),看梯度是否是有限数值(没有NaN或Inf),然后随便套一个简单任务,比如让模型复制输入序列,观察loss有没有下降趋势。只要能过这三关,编码器的核心逻辑基本就没问题了。
6.3 从编码器到完整Transformer:接下来还能做哪些扩展
这篇只写了编码器,但编码器理解之后,加上解码器就是完整的Transformer。解码器主要有两处扩展:
- 第一个自注意力子层要加causal mask,保证在生成当前位置时只看得到当前位置之前的token,不能偷看未来。
- 中间插入一个Cross-Attention子层,用解码器的Q去查询编码器输出的K和V,这样解码器才能从源语言表示里提取信息。
如果你接下来想做机器翻译、文本生成、或者干脆自己实现一个GPT风格的模型,都可以在这套编码器的基础上继续搭。而且这一步做扎实之后,再看BERT、ViT这些基于编码器的变体模型,思路就清晰很多——它们做的事情,本质上都是“如何把输入表示成更好的上下文向量”。
我自己的习惯是:每搭完一个模块,先用随机输入跑一遍,再单独写一个极小的单元测试,确认它的逻辑没毛病才继续往上层堆。这个习惯看着笨,但真的能省掉大量联调时的排查时间。手写Transformer编码器这件事,最大的价值不是让你以后不调包,而是当你真遇到效果不对、需要改结构的时候,心里确实知道每一行代码是在干什么,而不是靠试错瞎猜。