公司动态
手写Transformer基础组件:Linear、Embedding、初始化与反向传播
如果你只是把 Transformer 当成一行torch.nn.Transformer调用那你大概率会在论文复现、梯度调试、显存优化这类真实场景里卡住。很多看似诡异的问题比如权重为什么突然炸了、Embedding 的梯度为什么一直是 0、loss 为什么训到一半变 NaN最终都要回到前向计算图和反向传播的细节里找答案。这篇文章顺着斯坦福 CS 336 课程的 3.3 节思路把 Transformer 最基础的四个部分完整手写一遍Linear Layer、Embedding Layer、参数初始化、反向传播。为了保证每一步都透明可用代码用纯 NumPy 实现不依赖 PyTorch 的自动求导。读完你会得到一个能跑通梯度检查、能训练的最小 Python 实现也能真正理解nn.Linear和nn.Embedding底层到底做了什么。1. 为什么值得手搓 Transformer 的基础组件先给一个明确判断手写这些组件不是重复造轮子而是去理解深度学习框架隐藏的约定。举个例子。nn.Embedding接收整数索引很多人以为它只是一个查表操作。但它的反向传播并不是查表的逆过程而是scatter_add。如果你没写过这一步就不容易理解为什么 embedding 的梯度会累加到同一个 token 对应的行上也不容易解释多序列里重复 token 对参数更新的影响。再比如nn.LinearPyTorch 里的 weight 形状是(out_features, in_features)前向公式是y x W.T b。如果你只记住了y xW反向推导出来的dW形状一定是错的。这些细节在调包时被隐藏了但一旦你开始修改网络结构、做梯度裁剪、写自定义算子就必须把它们重新挖出来。适合读这篇文章的人有三类用 PyTorch 但没看过源码想搞懂常见层底层逻辑的开发者在理论课或面试中需要手推 Transformer 反向传播的学习者需要在没有自动求导的环境里实现模型比如 C 推理引擎、FPGA 部署的工程师。不适合读的人也有只关心 API 调用、不打算动底层逻辑的同学这篇文章对你来说信息密度偏高。2. 先建立计算图思维前向传播与反向传播手写反向传播唯一需要掌握的数学工具就是链式法则难的是追踪每一层张量的 shape 变化。一个完整的训练步包含两个阶段前向传播从输入到输出逐步计算并缓存中间结果。缓存的目的是给反向传播使用。反向传播从损失函数开始逐层计算梯度并把梯度沿着计算图回传到每个参数。这里最重要的工程习惯是每一层都要同时实现forward和backward并且backward只接收上游传过来的梯度dout返回当前层输入对应的梯度dx。参数自身的梯度如dW、db则缓存在层内部由优化器在之后统一更新。以最简单的标量传播为例z w * x loss z^2前向保存z和x。反向时dloss / dz 2 * zdloss / dx (dloss / dz) * (dz / dx) 2z * wdloss / dw (dloss / dz) * (dz / dw) 2z * x张量版本的推导完全一样只是乘法变成了矩阵乘法还多了转置。下面我们会用 Linear Layer 完整走一遍。3. 手写 Linear Layer全连接层的前向与反向3.1 前向公式与 shape 约定PyTorch 中nn.Linear的约定是输入x:(batch, in_features)权重W:(out_features, in_features)偏置b:(out_features,)广播后为(1, out_features)输出y x W.T b:(batch, out_features)注意这里W.T的形状是(in_features, out_features)所以x W.T能对齐。很多初学推导时把权重写成(in, out)导致反向推导出一堆转置很痛苦。3.2 反向传播推导设损失为L上游已经传回dout dL / dy形状为(batch, out_features)。对权重W形状必须是(out_features, in_features)dL / dW dout.T x维度验证(out, batch) (batch, in) (out, in)正确。对偏置bdL / db dout.sum(axis0, keepdimsTrue)因为偏置是逐元素相加所以梯度是dout在 batch 维上的求和。对输入xdL / dx dout W维度验证(batch, out) (out, in) (batch, in)正确。3.3 代码实现# 文件layers.py import numpy as np class Linear: 全连接层y x W.T b weight shape: (out_features, in_features) def __init__(self, in_features, out_features, init_methodxavier): if init_method xavier: self.W xavier_uniform((out_features, in_features)) elif init_method kaiming: self.W kaiming_uniform((out_features, in_features)) self.b np.zeros((1, out_features), dtypenp.float32) self.dW None self.db None def forward(self, x): # 缓存输入反向传播时使用 self.cache_x x return x self.W.T self.b def backward(self, dout): self.dW dout.T self.cache_x self.db dout.sum(axis0, keepdimsTrue) dx dout self.W return dx代码里的self.cache_x是计算图的关键。没有它反向传播无法计算dW。一个容易踩的坑是复用模块同一个 Linear 实例如果在前向中被调用了两次cache_x只会保存最后一次的输入梯度更新会出错。所以在实际深度学习框架里每个模块实例只被调用一次是基本约束如果你要在循环里共享参数必须手动累积梯度。4. 手写 Embedding Layer查表与梯度回传4.1 前向就是查表Embedding 层的权重是一个(num_embeddings, embedding_dim)的矩阵输入是整数 token id输出是对应行向量。前向非常简单output W[indices]这其实是一个 gather 操作。假设W是 4 行 3 列W [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6], [0.7, 0.8, 0.9], [1.0, 1.1, 1.2]] indices [2, 0, 2]那么输出是[[0.7, 0.8, 0.9], [0.1, 0.2, 0.3], [0.7, 0.8, 0.9]]注意索引2出现了两次这会在反向传播中产生梯度累加。4.2 反向是 scatter_add不是简单的查表逆运算对 embedding 来说输入是整数索引本身没有可学习的梯度所以反向传播只需要计算参数矩阵的梯度grad_W zeros_like(W) for each position i: grad_W[indices[i]] dout[i]这里必须使用累加因为同一个 token 可能在一条序列中出现多次也可能在一个 batch 的多条序列中出现。如果只做赋值梯度会丢失。NumPy 中有一个非常容易踩的坑grad_W[indices] dout实际上不是累加而是只在最后一次索引处赋值。必须使用np.add.at才能实现真正的原地累加。4.3 代码实现# 文件layers.py class Embedding: 嵌入层输入整数 token id输出对应行向量。 def __init__(self, num_embeddings, embedding_dim, std0.02): # Transformer 中常用小标准差高斯初始化 self.W np.random.normal(0.0, std, (num_embeddings, embedding_dim)).astype(np.float32) self.grad_W np.zeros_like(self.W) def forward(self, indices): # indices 形状任意元素为整数 self.cache_indices indices return self.W[indices] def backward(self, dout): # dout 形状与 W[indices] 一致 self.grad_W np.zeros_like(self.W) np.add.at(self.grad_W, self.cache_indices, dout) # 输入是整数索引不需要继续向上游传梯度 return None为什么np.add.at那么重要演示一下区别。假设indices [2, 0, 2]dout [0.1, 0.2, 0.3]普通赋值grad_W[indices] dout结果中第 2 行只落入了 0.3正确累加np.add.at(grad_W, indices, dout)第 2 行等于0.1 0.3 0.4。这一个细节往往是 Embedding 训练效果不稳定的隐藏原因之一。5. 参数初始化最容易忽略却最致命的步骤5.1 为什么初始化这么重要反向传播决定了梯度往哪个方向走但初始参数决定了模型从哪个点出发。初始化不当会直接引发两类问题参数全部相同比如全零神经元对称梯度一样网络无法学到不同特征参数方差过大深层网络中的信号逐层放大导致梯度爆炸loss 变成 NaN参数方差过小信号逐层衰减梯度消失loss 几乎不动。Transformer 对初始化尤其敏感因为自注意力中的 Softmax 会把数值差异指数级放大。5.2 三种常见初始化方法Xavier / Glorot 初始化适用于 tanh、sigmoid 这类对称激活函数。权重从均匀分布中采样bound gain * sqrt(6 / (fan_in fan_out)) W ~ Uniform(-bound, bound)Kaiming / He 初始化适用于 ReLU 类激活函数。它考虑到了 ReLU 会把一半神经元置零所以方差比 Xavier 略大gain sqrt(2 / (1 negative_slope^2)) bound gain * sqrt(3 / fan_in) W ~ Uniform(-bound, bound)小标准差高斯初始化GPT、BERT 等预训练模型常用N(0, 0.02)初始化 embedding 和大部分线性层。这是工程经验的产物配合残差连接和 LayerNorm 效果稳定。5.3 代码实现# 文件init.py import numpy as np def xavier_uniform(shape, gain1.0): Xavier / Glorot 均匀初始化适配 tanh / sigmoid 激活函数。 shape 约定为 (fan_out, fan_in)即权重矩阵完整形状。 fan_in, fan_out shape[1], shape[0] a gain * np.sqrt(6.0 / (fan_in fan_out)) return np.random.uniform(-a, a, sizeshape).astype(np.float32) def kaiming_uniform(shape, negative_slope0.0): Kaiming / He 均匀初始化适配 ReLU 激活函数。 negative_slope 是 LeakyReLU 的负斜率普通 ReLU 用 0。 fan_in shape[1] gain np.sqrt(2.0 / (1 negative_slope**2)) bound gain * np.sqrt(3.0 / fan_in) return np.random.uniform(-bound, bound, sizeshape).astype(np.float32) def normal_init(shape, std0.02): 小标准差高斯初始化Transformer 预训练模型中常用。 return np.random.normal(0.0, std, sizeshape).astype(np.float32)5.4 初始化选择的经验判断在实际项目中可以按下面的规则选型使用场景推荐初始化原因tanh / sigmoid 激活的全连接层Xavier保持前后向方差稳定ReLU 激活的全连接层Kaiming补偿 ReLU 造成的信号衰减Transformer 中的 EmbeddingN(0, 0.02) 或更小小标准差防止输出分布过散输出层分类头Xavier 或小正态避免初始 logits 过大击穿 Softmax一个关键提醒初始化方差需要和激活函数配套。把 Kaiming 初始化硬套到 tanh 上容易出现前几层输出饱和反向传播梯度接近 0。6. 串起来一个可运行的最小训练示例有了 Linear、Embedding 和初始化现在把它们组装成一个最小的可训练模型。这个模型的任务本身是玩具但足够验证前向、反向、参数更新整条链路。为了完整展示链式法则这里再补一个 LayerNorm 实现。它在 Transformer 中负责稳定分布反向传播公式也需要手推是很好的锻炼材料。# 文件layers.py class LayerNorm: Layer Normalization对最后一维做标准化。 def __init__(self, hidden_size, eps1e-5): self.gamma np.ones(hidden_size, dtypenp.float32) self.beta np.zeros(hidden_size, dtypenp.float32) self.eps eps def forward(self, x): self.x x self.mean x.mean(axis-1, keepdimsTrue) self.var x.var(axis-1, keepdimsTrue) # 总体方差除以 N self.x_hat (x - self.mean) / np.sqrt(self.var self.eps) return self.gamma * self.x_hat self.beta def backward(self, dout): N self.x.shape[-1] dx_hat dout * self.gamma x_hat self.x_hat dx 1.0 / np.sqrt(self.var self.eps) * ( dx_hat - np.mean(dx_hat, axis-1, keepdimsTrue) - x_hat * np.mean(dx_hat * x_hat, axis-1, keepdimsTrue) ) self.dgamma np.sum(dout * x_hat, axis0) self.dbeta np.sum(dout, axis0) return dx接着写 Softmax CrossEntropy。这里直接用数值稳定的 log-sum-exp 形式反向传播用probs - one_hot的经典技巧避免逐项推导 Softmax Jacobian。# 文件losses.py import numpy as np def softmax_cross_entropy(logits, label): logits: (1, num_classes), label: int 返回 (loss, probs) shifted logits - logits.max(axis-1, keepdimsTrue) exp_logits np.exp(shifted) probs exp_logits / exp_logits.sum(axis-1, keepdimsTrue) loss -np.log(probs[0, label] 1e-12) return loss, probs def softmax_cross_entropy_backward(logits, label): 返回 dlogits形状与 logits 相同。 _, probs softmax_cross_entropy(logits, label) dlogits probs.copy() dlogits[0, label] - 1.0 return dlogits最后是训练脚本# 文件train_mini.py import numpy as np from layers import Embedding, Linear from losses import softmax_cross_entropy, softmax_cross_entropy_backward np.random.seed(42) vocab_size 100 embed_dim 32 hidden_dim 64 num_classes 10 seq_len 3 embedding Embedding(vocab_size, embed_dim, std0.02) linear1 Linear(embed_dim, hidden_dim, init_methodkaiming) linear2 Linear(hidden_dim, num_classes, init_methodxavier) learning_rate 0.05 def forward(input_ids, label): emb embedding.forward(input_ids) # (seq_len, embed_dim) pooled emb.mean(axis0, keepdimsTrue) # (1, embed_dim) h np.maximum(linear1.forward(pooled), 0.0) # ReLU logits linear2.forward(h) # (1, num_classes) loss, probs softmax_cross_entropy(logits, label) return loss, probs, pooled, h, logits def backward(probs, pooled, h, logits, label): dlogits softmax_cross_entropy_backward(logits, label) # (1, num_classes) dh linear2.backward(dlogits) # (1, hidden_dim) drelu dh * (h 0.0) # ReLU 反向 d_pooled linear1.backward(drelu) # (1, embed_dim) # mean pooling 反向梯度均摊到每个 token d_emb np.broadcast_to(d_pooled / seq_len, (seq_len, embed_dim)) embedding.backward(d_emb) def update_params(): linear2.W - learning_rate * linear2.dW linear2.b - learning_rate * linear2.db linear1.W - learning_rate * linear1.dW linear1.b - learning_rate * linear1.db embedding.W - learning_rate * embedding.grad_W for step in range(100): input_ids np.random.randint(0, vocab_size, size(seq_len,)) label input_ids[0] % num_classes # 玩具任务用第一个 token 决定类别 loss, probs, pooled, h, logits forward(input_ids, label) backward(probs, pooled, h, logits, label) update_params() if step % 20 0: print(fstep {step}, loss: {loss:.4f})运行后可以观察到 loss 从约 2.3 逐步下降。由于任务是随机构造的loss 不会降到 0但稳定的下降趋势说明反向传播和参数更新链路是通的。这个最小示例的重点不是任务精度而是展示一条完整的因果链loss - dlogits - linear2 梯度 - ReLU 反向 - linear1 梯度 - mean pooling 反向 - embedding 梯度任何一环的公式或 shape 写错都会在最终 loss 上体现出来。7. 用梯度检查验证反向传播的正确性手写反向传播最容易出错而且错误往往很隐蔽loss 可能下降但梯度方向和大小并不正确。严谨的做法是用数值梯度中心差分和解析梯度做对比。数值梯度的公式g_numeric (f(x eps) - f(x - eps)) / (2 * eps)相对误差rel_error |g_analytic - g_numeric| / (|g_analytic| |g_numeric|)工程经验值是相对误差小于1e-6时可以认为解析梯度正确。# 文件grad_check.py import numpy as np def numerical_grad(fn, x, eps1e-5): 中心差分计算数值梯度。fn 接收 x返回标量。 x x.astype(np.float64) g np.zeros_like(x) it np.nditer(x, flags[multi_index]) while not it.finished: idx it.multi_index old x[idx] x[idx] old eps f_plus fn(x) x[idx] old - eps f_minus fn(x) x[idx] old g[idx] (f_plus - f_minus) / (2.0 * eps) it.iternext() return g def rel_error(analytic, numeric): denom np.maximum(np.abs(analytic) np.abs(numeric), 1e-8) return np.abs(analytic - numeric) / denom使用方式把某一层的参数当作输入构造一个从该参数到 loss 的函数然后对比linear.dW和numerical_grad的结果。# 文件check_linear.py from layers import Linear from grad_check import numerical_grad, rel_error linear Linear(4, 5, init_methodxavier) x np.random.randn(3, 4).astype(np.float32) # 前向 out linear.forward(x) # 模拟一个上游梯度 dout np.random.randn(3, 5).astype(np.float32) dx linear.backward(dout) # 以 loss sum(dout * out) 作为标量目标 def loss_from_W(W_flat): linear.W W_flat.reshape(5, 4) out linear.forward(x) return np.sum(dout * out) numeric_dW numerical_grad(loss_from_W, linear.W.astype(np.float64)) err rel_error(linear.dW, numeric_dW) print(W 相对误差:, err)如果误差在1e-6附近说明dW推导和实现正确。同理可以检查dx、db、Embedding 梯度和 LayerNorm 梯度。这里有个注意事项数值梯度计算量大每个参数维度都要做两次前向所以只适合对小型模块做验证不适合整网调试。8. 常见问题与排查方法问题现象可能原因排查方式解决方案loss 不下降或下降极慢初始化方差过小梯度消失打印梯度范数看是否接近 0换成 Kaiming 或增大学习率loss 为 NaN初始化方差过大或学习率过大打印 logits 和 softmax 输入改用 log-sum-exp 稳定 Softmax降低学习率Embedding 梯度不正确使用了grad_W[indices] dout在重复索引场景下对比期望值改用np.add.atgradient check 相对误差大反向公式推导错误或 shape 不一致单独检查每一层的dx和dW从输出层向输入层逐层排查同一个层被多次调用导致梯度错乱前向缓存被覆盖确认模块是否被复用了每次前向后立即反向或使用不同实例权重更新后模型不稳定偏置初始化太大检查偏置是否参与了初始化偏置统一初始化为 0最值得强调的是梯度检查当你的反向传播出现 bug 时不要靠loss 在下降来确认正确性。loss 下降只说明更新方向大体不差但可能梯度已经偏差很多导致模型收敛极慢或泛化差。任何新实现的反向传播都应该先用有限差分验证。9. 工程实践建议与后续方向9.1 保持 shape 显式化手写模块时建议在注释里写明每个中间张量的 shape。一旦 shape 不匹配错误会暴露得很快。例如# x: (batch, seq_len, hidden) # attn_weights: (batch, num_heads, seq_len, seq_len)9.2 写模块时就把梯度检查作为标配每写完一个模块立刻写一个针对它的梯度检查脚本。这比写完整个模型再调试高效得多。可以参考第 7 节的numerical_grad和rel_error工具。9.3 数值稳定性优先Softmax 必须做减最大值变换LayerNorm 的 epsilon 不能省略CrossEntropy 的 log 输入要加微小下界。这些看起来不起眼的操作在 Transformer 深层结构中能避免大量 NaN 问题。9.4 参数管理要统一手写实现中建议把参数和梯度放在同一个模块对象中而不是分散在全局变量里。等模块增多后最好实现一个简单的参数收集接口方便后续接入优化器。def parameters(self): 返回 (param, grad) 对列表。 return [(self.W, self.dW), (self.b, self.db)]9.5 后续学习路径这篇文章用手写方式覆盖了 Linear、Embedding、参数初始化和反向传播四个基础点。要继续深入建议按顺序推进实现 Multi-Head Self-Attention 的前向与反向重点推导Q、K、V三条梯度通路实现 Transformer Block组合 Attention、LayerNorm、MLP 和残差连接自己写一个最小 GPT 训练脚本在小型语料上观察 loss 变化对比 PyTorch 的autograd结果验证手写实现与框架一致性。到了第 2 步你会更容易理解深度学习框架为什么采用计算图 自动微分架构也会更清楚 Transformer 为什么需要残差连接和 LayerNorm。从零手写这些组件短期看是重复造轮子长期看是在积累排障直觉。下一次遇到梯度异常、Embedding 参数不更新、初始化导致训练不稳定这些问题时你会有明确的排查路径而不是靠运气调参。建议把这篇文章收藏起来动手写代码时对照使用。