手写Transformer核心模块:从零实现Linear、Embedding与反向传播

发布时间:2026/8/30 19:12:51
手写Transformer核心模块:从零实现Linear、Embedding与反向传播 这次我们来看一项很多人学完 Transformer 但依然没完全搞懂的事情不只是会调用nn.Transformer而是亲手把 Linear Layer、Embedding Layer、参数初始化和反向传播全部实现出来。本文对应的场景是 Stanford CS 336 3.3也就是在课程作业中要求用最朴素的 PyTorch 张量操作搭建语言模型组件并通过梯度检查验证手写反向传播的正确性。标题里的四个关键词其实就对应了 Transformer 落地中四个最容易翻车的点参数初始化决定模型能不能收敛反向传播决定你写的网络是不是真的“能学”Linear Layer 是所有前馈计算的地基Embedding Layer 则承担着离散 token 到连续向量的映射。如果这几个模块只是靠nn.Linear和nn.Embedding一带而过那遇到 loss 不降、梯度 NaN、维度对不上这类问题的时候往往只能瞎猜。自己手拼一遍之后很多报错一眼就能看出问题。这篇文章会从零实现一个可训练的最小 Transformer包含前向传播、手写 backward、参数初始化、梯度验证最后在一个小模型上完成端到端训练测试。全程使用 CPU 即可运行不依赖 GPU也不需要下载任何预训练权重。适合正在啃 Transformer 源码、准备算法面试、或者做课程作业时想把原理彻底吃透的读者。1. 核心能力速览先把这篇文章涉及的技术目标和运行约束讲清楚。下面的表格不是项目 README而是你读完这篇文章后应该具备的“能力验收清单”。能力项说明项目类型从零手写 Transformer 核心模块的教学实现来源Stanford CS 336 Language Modeling 课程 3.3 小节主题核心模块Linear Layer、Embedding Layer、LayerNorm、参数初始化、手写反向传播技术栈Python 3.9、PyTorch 2.x、NumPy 可选运行方式本地脚本运行支持 Jupyter Notebook硬件要求CPU 可跑显存不是必需是否依赖 nn.Linear不依赖手动实现 forward/backward是否依赖 nn.Embedding不依赖手动实现查表与梯度回传反向传播方式手动推导并实现通过 torch.autograd.gradcheck 验证验证方式形状测试、梯度对比、小模型端到端 loss 下降适合人群学习 Transformer 原理、面试准备、课程作业、源码阅读从表格可以看出这是一篇“先别用库、先把原理跑通”的教程。它的核心产出不是一个大而全的训练框架而是一套你可以自己掌控每个张量形状和每个梯度去向的最小实现。2. 适用场景与使用边界手写 Transformer 并不是为了让 PyTorch 用户抛弃官方 API。相反它适合在以下几个场景里发挥作用第一课程作业或实验报告。CS 336 这类课程的核心目标就是让学生脱离高级封装理解语言模型内部到底发生了什么。自己实现 Linear、Embedding、LayerNorm 之后再去看nn.TransformerEncoder的实现会明显感觉到“图层”不再是黑盒。第二算法面试准备。很多公司面试会问“Transformer 参数初始化为什么用 Xavier”“Embedding 梯度怎么回传”“softmax 反向传播的雅可比是什么”。这些问题的答案不亲手写一遍很难讲得清楚。第三源码阅读和二次开发。当你想修改某一层的行为比如给 Attention 换成线性注意力或者把 LayerNorm 换成 RMSNorm如果只依赖高层 API往往不知道改哪里但手写过底层模块之后你能很快定位到具体实现位置。同时也要明确使用边界。这套手写代码主要用于教学和原理验证不推荐直接用于大规模训练。手写循环、GPU 算子融合、FlashAttention 这些工程优化都不在范围内。如果目标是在生产环境训练大模型应该直接使用成熟框架PyTorch 的nn.Transformer、Hugging Face 的transformers或者 Megatron-LM 等。另一个必须强调的是合规边界。课程代码、作业要求、课件内容都有各自的版权和使用协议本文只讲解通用技术原理和常见实践不复制任何课程原始材料。训练时不要使用未经授权的数据、人脸信息、版权文本或用户隐私数据尽量在本地构造的玩具数据集上验证。3. 环境准备与前置条件手写层不需要复杂环境。建议使用 Python 3.9 以上版本和 PyTorch 2.x安装命令pip install torch numpy tqdm如果没有 GPU直接用 CPU 跑小模型即可。这里的关键不是训练速度而是验证前向输出形状与反向传播梯度是否正确。建议的工程目录结构transformer_from_scratch/ ├── layers/ │ ├── __init__.py │ ├── linear.py │ ├── embedding.py │ ├── layernorm.py │ └── attention.py ├── init.py ├── model.py ├── train.py └── tests/ └── test_gradcheck.py这种按层拆分的结构能让你在调试时只关心单个模块而不会一次面对整个网络。动手之前还需要确认三块数学基础矩阵乘法求导公式、链式法则、softmax 的雅可比矩阵。不需要研究得很深但至少要知道对于一个线性变换y x W从上层传回的梯度dy应该分别回传到dx和dW公式是dx dy W.T、dW x.T dy。后面所有代码都建立在这条基本规则上。4. 从零实现核心模块4.1 手写 Linear LayerLinear Layer 是最基础的计算单元公式很简单y x W b。但参数初始化和梯度回传是两个容易被忽略的细节。第一个细节是参数初始化。下面用 Xavier 均匀初始化它的上下界是bound sqrt(6 / (fan_in fan_out))。这种初始化适合线性层和 tanh 激活函数能让每层输出的方差保持在一个合理范围内避免深层网络时梯度消失或爆炸。第二个细节是反向传播。除了dx要回传给上一层还要计算dW和db供优化器更新。这里有一个常见错误把dW写成了grad_output.T x正确写法是x.T grad_output。推导方法是对损失函数求 W 的偏导保持张量形状对齐即可。import torch import math class LinearLayer: def __init__(self, in_features, out_features, biasTrue): self.in_features in_features self.out_features out_features self.bias bias # Xavier 均匀初始化 bound math.sqrt(6.0 / (in_features out_features)) self.W torch.empty(in_features, out_features).uniform_(-bound, bound) self.b torch.zeros(out_features) if bias else None def forward(self, x): self.x x out x self.W if self.b is not None: out out self.b return out def backward(self, grad_output): dx grad_output self.W.T dW self.x.T grad_output db grad_output.sum(dim0) if self.b is not None else None return dx, dW, db这段代码最值得注意的地方是forward里保存了self.x。这是手写反向传播的基础backward 需要用到前向时输入的原始值。很多新手写的反向传播出错不是因为公式推导错而是因为前向没有保存必要的中间变量。4.2 手写 Embedding LayerEmbedding Layer 本质上是一个查表操作。输入是 token id 序列输出是形状为[batch_size, seq_len, embedding_dim]的向量序列。它的反向传播很有特点并不是所有 embedding 行都会被更新只有当前 batch 中实际被查到的 token 行才有梯度。实现方式是用torch.index_add_把梯度累加到对应的行上。class EmbeddingLayer: def __init__(self, num_embeddings, embedding_dim, init_scale0.1): self.num_embeddings num_embeddings self.embedding_dim embedding_dim # 常用做法均匀分布 [-init_scale, init_scale] self.E torch.empty(num_embeddings, embedding_dim) self.E.uniform_(-init_scale, init_scale) def forward(self, tokens): self.tokens tokens return self.E[tokens] def backward(self, grad_output): dE torch.zeros_like(self.E) # 把梯度累加到被查到的行上 dE.index_add_(0, self.tokens.flatten(), grad_output.flatten(0, 1)) return dE这里有一个工程细节index_add_是原地操作调用前必须确保dE是全零矩阵。如果同一个 token 在一句话中出现多次它的梯度会自然累加这也符合 embedding 反向传播的定义。你不需要为它手动去重。4.3 参数初始化决定 Transformer 能不能收敛的隐藏开关参数初始化在 PyTorch 使用里常常被忽略因为nn.Linear已经封装好了默认初始化。但手写模块时每层初始化都必须自己处理。先看原始 Transformer 论文中的做法embedding 层使用均值 0、方差 1 的正态分布初始化之后乘以sqrt(d_model)。也就是说如果 embedding 矩阵初始化为N(0, 1)那么实际查表输出要再乘一个sqrt(d_model)。这样做的目的是让输入到 Attention 的向量范数不要太小避免 softmax 之前的所有点积都挤到极小值附近。GPT 风格的大模型通常采用更小的初始化方差例如N(0, 0.02)。原因很简单模型层数越深残差分支积累的隐状态方差越大逐层用小方差初始化可以抑制这种累积。如果你在深层 Transformer 里发现训练初期 logits 出现 NaN或者 loss 长时间不下降首先排查初始化方差是否过大。下面给出一个简单的参数初始化工具函数def init_linear_weight(W, modexavier): if mode xavier: fan_in, fan_out W.shape bound math.sqrt(6.0 / (fan_in fan_out)) W.uniform_(-bound, bound) elif mode gpt: nn.init.normal_(W, std0.02) else: raise ValueError(fUnknown init mode: {mode}) def init_embedding(E, d_model, modenormal): if mode normal: E.normal_(mean0.0, std1.0) E.mul_(math.sqrt(d_model)) elif mode uniform: scale 0.1 E.uniform_(-scale, scale) else: raise ValueError(fUnknown init mode: {mode})从实践角度看训练一个小模型时建议 Linear 层优先用 Xavier 均匀初始化Embedding 层可以用uniform(-0.1, 0.1)这种比较保险的小范围初始化。如果训练不收敛再尝试 GPT 风格的N(0, 0.02)并配合 warmup 学习率调度。4.4 手写反向传播并用 gradcheck 验证手写反向传播最难的不是某一层的公式而是整个链路的梯度形状。一个效率很高的验证方式是用torch.autograd.gradcheck。它会用数值差分计算梯度再和你的手动 backward 梯度对比误差在阈值内就说明反向传播正确。from torch.autograd import gradcheck torch.manual_seed(0) linear LinearLayer(64, 128) x torch.randn(2, 32, 64, requires_gradTrue, dtypetorch.float64) # 注意gradcheck 需要高精度通常使用 float64 linear.W linear.W.double() linear.b linear.b.double() output linear.forward(x) grad_output torch.randn_like(output) dx_manual, dW_manual, db_manual linear.backward(grad_output) # 用 autograd 计算参考梯度 output_double x.double() linear.W.double() linear.b.double() dx_auto torch.autograd.grad(output_double.sum(), x)[0] print(dx 误差:, (dx_manual - dx_auto).abs().max().item())gradcheck函数更严格可以直接传一个自定义函数def linear_forward_wrapper(W, b, x): return x W b linear_layer LinearLayer(16, 32) linear_layer.W linear_layer.W.double() linear_layer.b linear_layer.b.double() x torch.randn(4, 16, dtypetorch.float64, requires_gradTrue) W linear_layer.W.clone().requires_grad_(True) b linear_layer.b.clone().requires_grad_(True) check gradcheck( lambda w, b_: linear_forward_wrapper(w, b_, x), (W, b), eps1e-6, atol1e-5, ) print(gradcheck passed:, check)如果你是在课程作业里验证反向传播务必把所有参数和输入都转成float64。数值差分在float32下误差会很大导致明明公式正确gradcheck却报错。4.5 LayerNorm 与残差Transformer 稳定训练的第二道防线LayerNorm 在 Transformer 中负责把每一层的输出拉回稳定的数值范围。手写时可以直接使用简化形式y (x - mean) / sqrt(var eps)反向传播用一个更紧凑的公式class LayerNorm: def __init__(self, dim, eps1e-5): self.gamma torch.ones(dim) self.beta torch.zeros(dim) self.eps eps def forward(self, x): self.x x self.mean x.mean(dim-1, keepdimTrue) self.var x.var(dim-1, unbiasedFalse, keepdimTrue) self.x_hat (x - self.mean) / torch.sqrt(self.var self.eps) return self.gamma * self.x_hat self.beta def backward(self, grad_output): N self.x.shape[-1] dy grad_output * self.gamma dx (dy - dy.mean(dim-1, keepdimTrue) - self.x_hat * (dy * self.x_hat).mean(dim-1, keepdimTrue) ) / torch.sqrt(self.var self.eps) dgamma (grad_output * self.x_hat).sum(dim(0, 1)) dbeta grad_output.sum(dim(0, 1)) return dx, dgamma, dbeta这个反向公式的推导过程相当经典。核心思路是归一化操作同时依赖mean和var所以dx中必须包含dy.mean以及dy * x_hat的均值项。很多人手写 LayerNorm 反向时少了一项导致gradcheck一直不过几乎都是这个原因。残差连接则是把输入和子层输出相加它的反向传播不需要任何额外计算梯度在相加处直接“分叉”。实现时只要在 forward 里保存原始输入即可。4.6 组装一个最小 Transformer Block有了 Linear、Embedding、LayerNorm 和残差可以组装一个最小 Transformer Block。这里为了控制篇幅Attention 部分不手写 backward而是用 PyTorch 基础算子实现 forward反向交给 autograd。import torch import torch.nn.functional as F class MultiHeadAttention: def __init__(self, d_model, n_heads): self.d_model d_model self.n_heads n_heads self.head_dim d_model // n_heads self.Wq torch.empty(d_model, d_model) self.Wk torch.empty(d_model, d_model) self.Wv torch.empty(d_model, d_model) self.Wo torch.empty(d_model, d_model) for w in [self.Wq, self.Wk, self.Wv, self.Wo]: nn.init.xavier_uniform_(w) def forward(self, x): B, T, D x.shape q x self.Wq k x self.Wk v x self.Wv q q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) k k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) v v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) scores q k.transpose(-2, -1) / (self.head_dim ** 0.5) attn torch.softmax(scores, dim-1) out attn v out out.transpose(1, 2).reshape(B, T, D) return out self.Wo这里的关键知识点是缩放因子。Attention 公式里scores q k.T / sqrt(d_head)不能省略。如果不做缩放点积结果会随着d_head增大而变大softmax 会进入饱和区梯度接近 0训练会非常慢。5. 功能测试与效果验证5.1 测试 Linear 前向形状第一个测试很简单确认输出形状正确。torch.manual_seed(0) linear LinearLayer(64, 128) x torch.randn(2, 32, 64) out linear.forward(x) print(Linear input shape:, x.shape) print(Linear output shape:, out.shape) # 期望输出 (2, 32, 128)如果这里形状不对问题几乎一定出在W的定义方式上。W应该是(in_features, out_features)这样x W才能把最后一维从 64 映射到 128。5.2 测试 Embedding 反向是否正确构造一个只有 0、1、2 三个 token 的序列验证 embedding 各行的梯度。torch.manual_seed(0) embed EmbeddingLayer(num_embeddings3, embedding_dim8) tokens torch.tensor([[0, 1], [0, 2]]) out embed.forward(tokens) grad_output torch.randn_like(out) dE embed.backward(grad_output) print(dE shape:, dE.shape) print(dE row 0:, dE[0])判断依据第 0 行梯度应该等于输入中 token 0 对应位置梯度的累加。如果第 1 行或第 2 行也出现了不该有的梯度说明index_add_的索引写错了。5.3 用 gradcheck 验证全部手写层强烈建议给每个手写层都写一个 gradcheck。以 LinearLayer 为例用前文包装函数方式验证 W 和 b 的梯度EmbeddingLayer 则验证 E 矩阵的梯度。def embedding_forward_wrapper(E, tokens): return E[tokens] E torch.randn(10, 16, dtypetorch.float64, requires_gradTrue) tokens torch.tensor([[1, 2, 3], [4, 5, 6]], dtypetorch.long) gradcheck_check gradcheck(embedding_forward_wrapper, (E, tokens)) print(Embedding gradcheck:, gradcheck_check)注意gradcheck要求输入张量至少有一个是requires_gradTrue且所有浮点张量都建议使用float64。整数 token 张量不需要梯度。5.4 端到端小模型训练观察 loss 是否下降这是最后一个验证能证明整个链路从参数初始化到反向传播都没有大问题。import torch import torch.optim as optim torch.manual_seed(0) vocab_size 64 d_model 32 n_heads 4 block MultiHeadAttention(d_modeld_model, n_headsn_heads) embed EmbeddingLayer(vocab_size, d_model) x torch.randint(0, vocab_size, (2, 8)) y torch.randint(0, vocab_size, (2, 8)) optim_params [ {params: [block.Wq, block.Wk, block.Wv, block.Wo]}, {params: [embed.E]}, ] optimizer optim.Adam(optim_params, lr1e-3) loss_fn torch.nn.CrossEntropyLoss() for step in range(20): h embed.forward(x) h block.forward(h) logits h embed.E.T loss loss_fn(logits.view(-1, vocab_size), y.view(-1)) optimizer.zero_grad() # 这里故意用 autograd 计算整条链路梯度因为 MultiHeadAttention 没有手写 backward loss.backward() optimizer.step() if step % 5 0: print(step, step, loss, loss.item())这段代码用到了一个取巧但合理的设计embedding 矩阵既负责查表也作为输出投影矩阵。这种“共享嵌入”的写法在 GPT 系列中很常见。更好的实现是单独定义一个输出 Linear 层但从教学角度看共享嵌入能显著减少参数。判断标准如果 loss 在 20 步内明显下降说明参数初始化、前向、反向、优化器调度全部工作正常如果 loss 原地不动优先检查学习率是否太大或者太小其次检查初始化方差。6. 从“手写层”到“接口化调用”批量实验与可复用设计这个项目没有 HTTP API也不需要提供 Web 服务。但从工程化角度看你可以把自定义层设计成统一的“前向 反向”接口方便批量跑实验。比如定义一个 BaseLayer 风格class BaseLayer: def forward(self, x): raise NotImplementedError def backward(self, grad_output): raise NotImplementedErrorLinearLayer、EmbeddingLayer、LayerNorm 都继承这个接口这样在训练循环里可以统一调度。批量实验最常见的场景是扫描不同初始化模式对收敛速度的影响。可以写一个小脚本init_modes [xavier, gpt, uniform] learning_rates [1e-4, 1e-3, 1e-2] for init_mode in init_modes: for lr in learning_rates: model build_transformer(init_modeinit_mode) final_loss train_short(model, lrlr) print(init_mode, lr, final_loss)当你把“参数初始化”和“反向传播”都变成可配置项之后就能系统性地观察为什么 GPT 风格的小方差初始化在深层模型中表现更好为什么 Xavier 在浅层模型中已经足够。这种实验往往比单纯阅读论文更有价值。如果你确实想把模块暴露成服务也可以额外包一层 FastAPI但这属于教学项目之外的扩展不推荐在原理学习阶段引入。7. 资源占用与性能观察手写模块的最大问题是性能远低于 PyTorch 原生算子。因为这里为了教学把 forward 和 backward 的中间张量都保存在内存中也没有做算子融合。如果是大型模型训练这种写法会占用大量显存。观察资源使用可以这样写import time import psutil x torch.randn(4, 128, 256) linear LinearLayer(256, 256) # warmup linear.forward(x) t0 time.time() for _ in range(10): out linear.forward(x) _, _, _ linear.backward(torch.randn_like(out)) t1 time.time() print(平均耗时:, (t1 - t0) / 10) mem psutil.Process().memory_info().rss / 1024 / 1024 print(当前内存占用 MB:, mem)如果训练时显存不足有几种通用降载方法减小 batch size减小序列长度降低 d_model 和 n_heads把不必要的中间变量从 forward 中移出例如 forward 结束后只保留反向需要用到的张量不保存完整 attention 矩阵对大模型训练改用半精度或混合精度。从性能观察的角度说手写版的价值不在于快而在于你能清楚看到每个操作的时间和空间开销。这也是理解后续 FlashAttention 为什么要合并算子、减少访存的起点。8. 常见问题与排查方法问题现象可能原因排查方式解决方案gradcheck 一直报错没有开启 float64检查输入张量和参数 dtype全部转成torch.float64Linear 反向维度报错dW写成grad_output.T x打印x.shape、grad_output.shape改为x.T grad_outputEmbedding 梯度全部为 0查表后没有保存 token 索引检查forward是否保存self.tokens在 forward 中保存输入 tokenloss 不下降学习率不合适或初始化方差太小打印梯度均值、方差调大 learning rate 或改用 GPT initloss 出现 NaN初始化方差太大或 Attention 没有缩放打印 logits 是否存在 inf使用sqrt(d_head)缩放减小初始化方差LayerNorm 反向公式缺项忘记处理 mean 和 var 的依赖用 gradcheck 定位使用包含dy.mean与dy*x_hat项的简化公式Attention 维度 reshape 错误Q、K、V 的 view/transpose 顺序不对打印q.shape、k.shape、scores.shape先 view 成[B, T, n_heads, head_dim]再 transpose 成[B, n_heads, T, head_dim]CPU 训练太慢batch size / 序列长度偏大用 time 统计 forward/backward 耗时减小 batch size 和 seq_len从上表能看出手写 Transformer 的大部分报错都集中在维度形状和梯度公式上。我的建议是遇到问题不要直接搜答案先打印梯度的 shape 和数值范围往往能更快定位。9. 最佳实践与使用建议写手写 Transformer 项目时有几条通用的工程建议值得遵守。第一按层拆分文件。不要把所有自定义层堆在一个 Python 文件里。按linear.py、embedding.py、layernorm.py、attention.py拆分每层只做一件事调试时直接定位对应文件。第二每写一层就立即验证一层。不要等到所有模块写完了再验证。验证工具就是torch.autograd.gradcheck它比任何肉眼观察都可靠。每层验证通过之后再进入下一步。第三先跑最小配置。第一个实验建议使用vocab_size64、d_model32、n_heads4、batch_size2、seq_len8。这么小的配置可以在几秒内完成完整训练循环。如果最小配置都能跑通再逐步放大。第四保留一个 baseline 对照。用 PyTorch 官方nn.Linear、nn.Embedding或nn.TransformerEncoder实现同样的网络不断对比前向输出和梯度。这样能快速排除“手写层写错”和“训练逻辑写错”两类问题。第五模型文件、训练代码和测试代码分目录管理。避免把临时调试代码和正式代码混在一起。如果你要批量跑初始化实验建议加一个配置字典或者命令行参数。第六涉及真实数据时注意合规。本文教学项目只需要随机 token 或自己构造的玩具数据不需要真实文本。如果后续扩展项目需要真实语料务必确认数据来源合法、授权完整不涉及个人隐私和版权内容。10. 总结与下一步这个手写项目最值得验证的功能就是“参数初始化 手写反向传播”的组合。你先跑通一层 Linear 的 gradcheck再试一个最小 Transformer 的 loss 下降就会理解为什么Transformer的每一层都需要精心设计梯度通路。最容易踩的坑有三个Linear 反向里dW的维度顺序、LayerNorm 反向里缺少归一化依赖项、Attention 里的缩放因子。这三个坑只要踩过一个以后再看到类似报错就会很敏感。下一步建议做三件事第一给 MultiHeadAttention 也实现手写 backward完整覆盖所有模块第二加入 causal mask让 Attention 只能看到当前位置之前的 token变成一个真正的自回归语言模型第三把学到的参数初始化经验搬到一个小型文本数据集上观察 loss 曲线在不同 init 和 learning rate 下的表现。做完这三步你对 Transformer 的掌握就不再停留在“会用层”的层面而是真正进入了“能写层”的阶段。