复现一个 Transformer Encoder:从论文公式到最小 PyTorch

复现一个 Transformer Encoder:从论文公式到最小 PyTorch

系列:开源 AI 论文复现实验与代码解读

日期:2026-08-24

适合读者:研究生、科研新人、希望真正理解 Transformer 代码的工程读者

检索日期:2026-08-24

目录

  1. 为什么要从零复现 Encoder
  2. 原论文中的 Encoder 到底包含什么
  3. 从公式到张量形状
  4. 最小 PyTorch 实现怎么读
  5. 最小实验与验收标准
  6. 与 PyTorch 官方实现对照
  7. 最容易踩的复现陷阱
  8. 可以继续研究的问题
  9. 总结与参考资料

为什么要从零复现 Encoder

"会调用 Transformer"和"理解 Transformer"之间,差的往往不是更多公式,而是把公式、张量和代码逐项对齐的能力。研究生做论文复现时,真正消耗时间的通常不是写出 model(x),而是回答这些细节:输入是 [batch, length, hidden] 还是 [length, batch, hidden]?mask 中的 True 表示保留还是屏蔽?多头拆分后为什么要 transpose?残差分支加的是归一化前还是归一化后的表示?为什么输出形状正确,loss 却一直是 nan

从零实现一个 Encoder 的价值,是建立一份可检查的"结构账本"。每个模块都应有输入形状、输出形状、数学含义和不变量。这样再阅读 BERT、Vision Transformer、CLIP 或多模态模型时,看到的就不再是一组陌生类名,而是同一基本结构在位置编码、归一化、激活函数和注意力模式上的变体。

这也决定了本文的复现边界。原始 Transformer 是编码器-解码器机器翻译系统,论文的 WMT 结果依赖数据处理、训练预算、学习率调度、label smoothing、beam search 和 checkpoint averaging。一个本地小脚本不可能证明这些结果。我们只复现 Encoder 的计算图与关键不变量,把"大结果复现"拆成可验证的第一步。

原论文中的 Encoder 到底包含什么

2017 年的《Attention Is All You Need》用自注意力替代循环网络和卷积,让一个序列中的所有位置可以并行计算。原论文的 Encoder 由 N 个相同结构的层堆叠而成,每层包含两个子层:Multi-Head Self-Attention 和 Position-wise Feed-Forward Network。每个子层外都有残差连接与 LayerNorm。

对输入 token 序列,最底部先查 embedding,再乘以 sqrt(d_model),随后加上同维度的位置编码。位置编码是必要的,因为自注意力本身只根据内容计算 token 间关系;如果交换两个 token,同时交换对应的 Q、K、V,纯注意力不会自动知道顺序发生了变化。

原论文使用正弦和余弦函数构造固定位置编码:

P E ( p o s , 2 i ) = sin ⁡ ( p o s / 10000 2 i / d m o d e l ) PE(pos, 2i)=\sin\left(pos/10000^{2i/d_{model}}\right) PE(pos,2i)=sin(pos/100002i/dmodel)

P E ( p o s , 2 i + 1 ) = cos ⁡ ( p o s / 10000 2 i / d m o d e l ) PE(pos, 2i+1)=\cos\left(pos/10000^{2i/d_{model}}\right) PE(pos,2i+1)=cos(pos/100002i/dmodel)

偶数维和奇数维使用不同相位,不同维度对应不同频率。实现时最重要的不是背公式,而是确认 PE 的形状可广播到 [B,T,d_model],并作为 buffer 跟随模型移动设备,却不参与梯度更新。

Encoder 输出仍然是长度为 T 的序列表示,而不是一个分类结果。每个位置都已通过自注意力融合全序列信息。BERT 后来把这种双向 Encoder 表示用于预训练,再针对分类、问答等任务增加输出头;这也是"Encoder 不是只为机器翻译服务"的典型例子。

从公式到张量形状

1. Scaled Dot-Product Attention

注意力的核心公式是:

A t t e n t i o n ( Q , K , V ) = s o f t m a x ( Q K T d k + M ) V Attention(Q,K,V)=softmax\left(\frac{QK^T}{\sqrt{d_k}}+M\right)V Attention(Q,K,V)=softmax(dk QKT+M)V

若输入 x[B,T,d_model],先通过线性映射得到 Q、K、V。拆成 h 个头后,三者形状都变为 [B,h,T,d_head],其中 d_head=d_model/h。因此 d_model 必须能被头数整除。

Q @ K.transpose(-2,-1) 得到 [B,h,T,T]。前一个 T 表示"哪个 Query 正在提问",后一个 T 表示"它能查看哪些 Key"。除以 sqrt(d_head) 是为了控制点积方差;维度增大时,未经缩放的分数容易进入 softmax 的饱和区,梯度会变得不友好。

mask 应在 softmax 前作用于 score。对 padding mask,True 通常表示对应 Key 是填充位置,需把该列 score 设成一个极小值。这里尤其容易写反:屏蔽的是 Key 维,也就是 [B,1,1,T] 广播到所有头和所有 Query,而不是随手把输出位置清零。PyTorch 官方 MultiheadAttention 文档同样规定,二值 key_padding_mask 中的 True 表示忽略该 Key。

2. Multi-Head 的意义

单头注意力只在一个投影空间内计算关系。多头注意力让不同头拥有独立的 Q、K、V 子空间,再把各头结果拼回 [B,T,d_model],最后通过输出投影 W^O 混合。

M u l t i H e a d ( Q , K , V ) = C o n c a t ( h e a d 1 , ... , h e a d h ) W O MultiHead(Q,K,V)=Concat(head_1,\ldots,head_h)W^O MultiHead(Q,K,V)=Concat(head1,...,headh)WO

这不意味着每个头必然学出人类可命名的语言学关系。注意力权重可以用于诊断,但"某个头看向某个 token"不能直接等同于因果解释。下一篇的 Attention 可视化会专门讨论这一边界。

3. Position-wise FFN

前馈网络独立作用于每个位置:

F F N ( x ) = max ⁡ ( 0 , x W 1 + b 1 ) W 2 + b 2 FFN(x)=\max(0,xW_1+b_1)W_2+b_2 FFN(x)=max(0,xW1+b1)W2+b2

它先把 d_model 扩展到 d_ff,经过 ReLU,再投影回 d_model。注意力负责跨位置混合,FFN 负责对每个位置做非线性特征变换。两者分工不同,不能因为"attention is all you need"就删掉 FFN。

4. 残差与 LayerNorm

原论文写法是 LayerNorm(x + Sublayer(x)),即 Post-LN。PyTorch 的 TransformerEncoderLayer 默认 norm_first=False,也对应先做子层与残差,再归一化;设置 norm_first=True 则变为 Pre-LN。后续研究表明,两种顺序会改变深层 Transformer 的梯度行为和 warm-up 需求。因此复现时必须把归一化顺序视为实验变量,不能只看模块是否都存在。

最小 PyTorch 实现怎么读

配套脚本位于 code/minimal_transformer_encoder.py。建议不要从 main() 开始,而按以下路径阅读。

第一步看 MultiHeadSelfAttention.forward()。重点追踪 QKV 合并投影、viewpermute、score 矩阵、mask 广播、softmax、各头拼接这七个动作。只要其中一次维度交换错了,代码可能仍能运行,却把 batch、head 或 sequence 的语义混在一起。

第二步看 EncoderLayer.forward()。脚本同时保留 Post-LN 和 Pre-LN 两条路径,默认使用原论文式 Post-LN。这里应明确残差分支中的 x 指向哪一阶段,避免无意中把 norm(x) 与未归一化的残差相加。

第三步看 TransformerEncoder.forward()。它完成 embedding 缩放、位置编码和 N 层堆叠。这里用 ModuleList 创建独立层,而不是重复引用同一个 layer 对象;否则看起来堆叠了 N 层,实际却共享全部参数。

第四步看 run_invariant_checks()。它不检查"模型聪不聪明",只检查更基础、也更可靠的事实:输出是否为 [2,5,64],每层注意力是否为 [2,4,5,5],softmax 后每行是否求和为 1,被 padding mask 屏蔽的 Key 概率是否为 0。复现项目应先通过这些不变量,再谈训练曲线。

最小实验与验收标准

安装依赖后,可先只运行结构检查:

bash 复制代码
python3 -m pip install -r code/requirements.txt
python3 code/minimal_transformer_encoder.py --check-only

再运行合成任务:

bash 复制代码
python3 code/minimal_transformer_encoder.py --steps 80 --seed 7

任务随机生成长度为 12 的 token 序列,并平衡构造两类样本:正类的首 token 与尾 token 相同,负类则确保两者不同。分类头只读取第一个位置的 Encoder 表示,因此模型若要利用最后一个 token,必须通过层内信息交互把远端信息传回首位置。这个任务不是 benchmark,也不能证明实现等价于原论文训练系统;它只是一个便宜的端到端连通性检查。

建议把验收分成四层。

层级 检查内容 失败时优先排查
结构 所有张量形状与头数一致 viewpermuted_model % n_heads
数值 loss、梯度和 attention 无 nan/inf mask、学习率、初始化、全 padding 行
学习 固定 seed 后 loss 有下降趋势 标签构造、梯度是否更新、dropout、任务是否可学
对照 与官方层在相同配置下输出形状和 mask 语义一致 batch_first、Pre/Post-LN、激活、bias、dropout

本次实际验证结果是:脚本已通过 Python 语法编译;当前环境因缺少 torch,未执行结构断言和训练循环。安装依赖后,首先应运行 --check-only,再观察 toy loss,而不是直接引用一个未经运行的准确率。

与 PyTorch 官方实现对照

PyTorch 官方 TransformerEncoderLayer 明确由 self-attention 和 feed-forward network 构成,并提供 d_modelnheaddim_feedforwarddropoutactivationbatch_firstnorm_first 等参数。把本文实现与官方层逐项对照,可以得到一张最实用的阅读表。

本文模块 PyTorch 官方对应 需要核对的语义
MultiHeadSelfAttention nn.MultiheadAttention 输入布局、mask 中 True 的含义、是否返回每个 head 权重
EncoderLayer nn.TransformerEncoderLayer Post-LN / Pre-LN、激活函数、dropout 位置
ModuleList[layers] nn.TransformerEncoder 层数、参数是否独立、最终 norm
手写 attention 公式 scaled_dot_product_attention 数值稳定性、融合 kernel、mask 类型

手写版本适合学习和插桩,不适合直接替代生产实现。PyTorch 当前教程建议用 scaled_dot_product_attention、Nested Tensor、torch.compile() 和 FlexAttention 等底层组件构建更灵活高效的层。这里的差异不是"公式变了",而是官方实现还要处理 kernel 调度、不同 dtype、设备、稀疏长度、编译和性能快路径。

最容易踩的复现陷阱

第一,忘记 embedding 缩放或位置编码。小任务也许仍能拟合,但这已经不是原论文的输入构造,不能把结果当成同一实现。

第二,softmax 维度写错。attention 应沿最后一个 Key 维归一化。如果沿 head 或 Query 维做 softmax,输出形状仍完全正确,错误却很隐蔽。

第三,mask 语义混乱。不同库可能用 1 表示保留,也可能用 True 表示屏蔽。不要凭记忆转换;先读当前版本文档,再写一个两行样例验证被屏蔽位置的概率确实为 0。

第四,所有 Key 都被屏蔽。若某个样本的一行 score 全是极小值,softmax 可能产生无意义分布,部分实现或 dtype 下还可能出现 nan。数据管线应保证至少有一个有效 token,或显式定义空序列行为。

第五,混淆 Encoder 自注意力和 Decoder 因果注意力。Encoder 通常允许每个位置查看整个输入,不应默认加上三角 causal mask。把 Decoder 的"不能看未来"照搬到 BERT 风格 Encoder,会悄悄改变任务定义。

第六,层复制成参数共享。[layer] * N 只是把同一个 Python 对象引用 N 次。若论文没有明确共享参数,应为每层实例化独立模块。

第七,只报告最终 loss。一个可信复现还应记录 PyTorch 版本、设备、seed、数据生成方式、参数量、训练步数、优化器与提交版本。toy task 的成功只能说明实现具有基本学习能力,不能外推到翻译质量、长序列能力或真实语料性能。

可以继续研究的问题

  1. 在层数从 2 增加到 24 时,Post-LN 与 Pre-LN 的梯度范数、warm-up 需求和收敛速度如何变化?
  2. 固定参数量后,增加 head 数是否总能提升合成关系任务?当 d_head 太小时会发生什么?
  3. 正弦位置编码、可学习绝对位置编码、RoPE 与 ALiBi 在长度外推上的差异,如何用可控数据验证?
  4. 手写 attention 与 PyTorch scaled_dot_product_attention 在 CPU、CUDA、不同 dtype 下的数值误差与速度差异多大?
  5. padding mask、causal mask 和任意稀疏 mask 能否统一为一个经过单元测试的接口?
  6. 注意力权重与模型预测之间是否具有因果关系?遮蔽高权重连接后,输出变化是否真的更大?

这些问题都适合从本文脚本扩展。研究价值不在于再造一个更大的 Transformer,而在于把一个变量、一个假设和一个测量协议说清楚。

总结

复现 Transformer Encoder 的第一步,不是追求与原论文相同的 BLEU,而是建立公式、张量、代码和测试之间的对应关系:QKV 如何拆头,score 为什么除以 sqrt(d_head),mask 在哪个维度生效,注意力后为何还需要 FFN,残差与 LayerNorm 的顺序又如何成为可研究变量。

本文的最小实现刻意保留可读性,并用形状、概率和 mask 不变量约束正确性。它已经完成静态语法验证,但由于当前环境没有 PyTorch,运行时结果仍标注为待人工核验。这样的边界说明不是缺点,而是复现记录最重要的一部分:只报告真正做过的验证。

参考资料

检索日期:2026-08-24。PyTorch 文档会随版本变化;正式复现实验应同时记录本地 torch.__version__、设备与提交版本。

  1. Vaswani et al., Attention Is All You Need, 2017;arXiv 当前页面包含 2023 年修订版本。
  2. Harvard NLP, The Annotated Transformer,带注释的 PyTorch 实现与 notebook。
  3. PyTorch, TransformerEncoderLayer,当前官方 API 说明。
  4. PyTorch, MultiheadAttention,输入形状、key_padding_maskattn_mask 语义。
  5. PyTorch Tutorials, Transformer building blocksscaled_dot_product_attention、Nested Tensor 与自定义层的当前建议。
  6. Harvard NLP, the_annotated_transformer.py,位置编码、模型构造、初始化与训练代码路径。
  7. Devlin et al., BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding, 2018。
  8. Xiong et al., On Layer Normalization in the Transformer Architecture, ICML 2020。
相关推荐
June bug1 小时前
【HCIA- AI(正课)】2.2 神经网络组成
人工智能·深度学习·神经网络
TAN-90°-1 小时前
Deep Learning for Computer Vision——Recurrent Neural Networks
数据结构·人工智能·rnn·深度学习·神经网络·机器学习·计算机视觉
海兰1 小时前
【应用】Ubuntu 24 搭建大数据 & AI 应用云原生容器化实践
大数据·人工智能·ubuntu
ZGi.ai1 小时前
ZGI 父子分块:连接检索片段与完整上下文
人工智能·算法·知识库·企业ai·zgi·父子分块
yyywxk1 小时前
ICCV 2025 目标检测(object detection)方向上接收论文总结
人工智能·目标检测·计算机视觉
欧特克_Glodon1 小时前
OpenCV计算机视觉开发入门与实践<十七>:点运算与灰度变换概述
c++·人工智能·opencv·计算机视觉
chen_zn951 小时前
《VLA 系列》MemoryVLA++ | 感知-认知记忆 | 潜空间未来想象 | 论文与源码边界解析
人工智能·具身智能·vla
狂云歌1 小时前
AI时代,学什么,怎么学
人工智能·学习