最近在跟 CS 336 系列作业,到了 3.3 这一节,主题很集中:自己动手实现 Transformer 里最基础的 Linear Layer、Embedding Layer,并补齐参数初始化、反向传播和梯度更新。这节做完,你会明显感觉到,之前看的那些“Transformer 模型详解”和“手撕 Transformer”教程,几乎都建立在同样的矩阵运算和梯度回传逻辑上。如果你也想从零开始写一个能训练的小模型,这篇就把我踩过的关键步骤和判断标准拆开讲。
先给结论:这节作业真正值钱的地方,不是把 PyTorch 里现成的nn.Linear调出来用,而是用纯矩阵运算把前向和反向写出来,并且保证梯度数值是对的。很多人卡在三个地方:参数初始化选不好导致梯度消失或爆炸;反向传播公式推导时形状搞反;Embedding 层做反传时不知道梯度要累加到哪个位置。这些问题都可以通过小样例和梯度检查提前发现,不要一上来就训练大模型。
我建议先别急着跑完整 Transformer,而是把 Linear Layer 和 Embedding Layer 单独拿出来,用一个小批数据验证前向、反向和参数更新。下面按实际落地顺序拆一遍。
1. 先把作业里的关键模块拆清楚,再动手写代码
1.1 这节作业到底在做什么
CS 336 3.3 这一段,核心是完成一套从零手写的最小 Transformer 组件。从标题看,它覆盖了四个主题:参数初始化、反向传播、Linear Layer、Embedding Layer。翻译成实际任务就是:
- 实现 Linear Layer 的前向计算和反向计算;
- 实现 Embedding Layer 的前向查表和反向梯度累加;
- 用合适的参数初始化方法给权重赋初值;
- 把梯度传回去,更新参数,确认 Loss 能下降。
这四个主题互相关联。参数初始化影响反向传播的梯度尺度;Linear Layer 和 Embedding Layer 的梯度形状如果不匹配,后续整个网络都会出问题。
1.2 一前一后两条链路分别是什么
从数据角度理解,Transformer 内部有两条链路:
- 前向链路:输入 token ID -> Embedding 查表 -> 位置编码 -> 多层 Transformer Block -> 输出 logits -> Loss;
- 反向链路:Loss -> 输出层梯度 -> 每层 Linear 的权重梯度 -> 继续往前传 -> Embedding 表梯度 -> 更新参数。
作业里 3.3 重点处理的,是这两条链路里最底层的两个模块。你不需要一次写完整 Transformer,先把这两个模块的“前向怎么算、反向怎么传、参数怎么更新”彻底搞明白,后面的多头注意力、LayerNorm、残差连接都只是在这些基本运算之上叠加新操作。
1.3 建议的验证顺序
我建议把验证顺序拆成四步:
- 用一个很小的随机输入测试 Linear 前向,确认输出形状和数值;
- 用相同输入测试 Embedding 前向,确认查表结果正确;
- 手工构造一个标量 Loss,调用自己写的反向函数,查看每个梯度形状;
- 和 PyTorch 自动微分结果对比,差距在 1e-6 以内再继续。
这个顺序能帮你把问题隔离。如果形状和数值都对,再进入完整模型训练,否则后面一步错步步错。
2. 参数初始化不是玄学,它决定梯度能不能正常传播
2.1 Linear Layer 参数初始化为什么不能全零
很多人第一次手写网络时,会把权重初始化为 0。这在反向传播看来是灾难性的。假设某一层所有权重都是 0,那么同一层内每个神经元的输入完全一样,反向传播时梯度也完全一样,所有参数会同步更新,等于这一层只有一个有效神经元在表达,模型容量被严重浪费。
对 Linear Layer,常见做法是采用均匀分布或正态分布,尺度根据输入输出维度决定。比如 Xavier/Glorot 初始化:
import math # 以普通 Python 为例 def xavier_uniform(shape): fan_in, fan_out = shape[0], shape[1] if len(shape) > 1 else shape[0] limit = math.sqrt(6.0 / (fan_in + fan_out)) return [[random.uniform(-limit, limit) for _ in range(fan_out)] for _ in range(fan_in)]在实现时,可以用np.random.uniform或 PyTorch 的torch.empty配合uniform_填充。关键是理解公式里的fan_in和fan_out:
fan_in是这一层输入的特征数;fan_out是这一层输出的特征数;- 方差同时考虑两者,是为了让信号在前向和后向传播时尺度都比较稳定。
如果做的是深度 Transformer,很多实现会使用更精细的初始化,比如把某些权重的标准差除以2 * num_layers,用来配合残差连接的方差累积。这个细节可以后面再调,第一阶段先把基础初始化写对。
2.2 Embedding Layer 的初始化方式与常见选择
Embedding Layer 本质上是一张二维查找表,形状是vocab_size x embedding_dim,每一行对应一个 token 的向量表示。初始化方式通常也是随机初始化,常见选择有两种:
- 正态分布:
N(0, 0.02); - 均匀分布:
U(-0.1, 0.1)或U(-1/sqrt(embedding_dim), 1/sqrt(embedding_dim))。
在 GPT、BERT 这类模型里,N(0, 0.02)是一个很常见的取值,但它不是唯一标准。作业里如果只要求“用合理的初始化”,我一般会先用标准差较小的正态分布,避免一开始产生过大的输出信号。
需要注意一点:Embedding 层和 Unembedding 层(输出词表映射层)在部分实现中会共享权重。共享权重的好处是减少参数量,但坏处是初始化条件、梯度回传路径会变得复杂。CS 336 里如果先做不共享的版本,建议把两层分开定义,验证通过后再考虑共享。
2.3 初始化尺度对梯度的影响怎么判断
初始化尺度不是越小越好。假设权重初始值都在 1e-8 附近,前向输出会接近 0,反向梯度传到线性层时,虽然数值不大,但参数更新会非常慢;如果初始值太大,比如标准差为 1.0 且网络有 12 层,前向信号经过多层矩阵乘法后可能变成非常大的数,梯度爆炸风险很高。
一个直观判断方法:跑一次前向,看每一层输出均值和方差。如果第一层输出方差还在 1 左右,第 5 层已经变成 1e-10,说明初始化尺度偏小;如果变成 1e10,说明偏大。对这种问题,可以先调整初始化标准差,再考虑 LayerNorm。
注意:不要一上来就开完整训练,先用固定随机种子打印每层输出方差。初始化异常在训练中通常表现为 Loss 不下降或直接变 NaN,但这个现象往往滞后,不如前向检查来得快。
3. 手写 Linear Layer 与 Embedding Layer 的前向实现
3.1 Linear Layer 前向的矩阵形状核对
一个标准 Linear Layer 做的事情是:
output = input @ W.T + b这里input形状通常是(batch_size, seq_len, in_features)或(batch_size, in_features),W形状是(out_features, in_features),b形状是(out_features,)。
在实现时,最容易踩的坑是矩阵乘法方向。如果input @ W.T和W @ input.T搞混,输出形状会直接不对。建议写一个辅助函数打印每一步形状:
def linear_forward(x, weight, bias): # x: (batch_size, in_features) out = x @ weight.T + bias return out对于 Transformer 里的输入,通常需要先对(batch_size, seq_len, in_features)的输入做处理。可以把batch_size * seq_len当成一个维度来算,也可以直接用 PyTorch 的矩阵乘法自动处理。如果是纯 Python 实现,就要先 reshape 成二维矩阵,做完再 reshape 回去。
3.2 Embedding Layer 前向的查表实现
Embedding Layer 前向就是查表。输入是一组 token ID,输出是对应 ID 所在行的向量。比如embedding_table形状是(vocab_size, embedding_dim),输入token_ids = [2, 0, 1],输出就是把第 2 行、第 0 行、第 1 行拼出来。
在 PyTorch 里,可以直接用embedding_table[token_ids]实现查表。但如果要求手写反向传播,最好把这个“取行”的过程理解成矩阵运算:输入 one-hot 向量拼成的矩阵(batch, vocab_size),乘上embedding_table,会得到(batch, embedding_dim)。虽然实际实现不会真的构造 one-hot,但理解这种方式,对反向传播很有帮助。
3.3 用一个小样例验证前向结果
我建议使用固定随机种子构造一个非常小的任务:
vocab_size = 5embedding_dim = 4batch_size = 2seq_len = 3
然后手动计算期望输出,或者用 PyTorch 自动微分作为参照。这样跑完前向,你可以同时检查三件事:
- Embedding 输出的形状是不是
(2, 3, 4); - Linear 输出形状是不是
(2, 3, out_features); - 数值是否和直接用矩阵乘法的结果一致。
如果前向结果对不上,优先检查维度顺序和初始化的随机种子。很多时候不是逻辑错,而是weight定义成了(in_features, out_features),但代码里按weight.T算,结果就多了一次转置。
4. 反向传播:从 Loss 到梯度,再到参数更新
4.1 先画计算图,再写公式
手写反向传播,最容易犯的错误是跳过计算图直接套公式。对 Linear Layer 来说,前向可以拆成两步:
z = x @ W.T + bloss = f(z)
如果上游传过来的梯度是dz = dloss / dz,那么:
- 对
W的梯度:dW = dz.T @ x - 对
x的梯度:dx = dz @ W - 对
b的梯度:db = dz.sum(axis=0)
这里面的维度规则是:任何梯度的形状必须和原始变量形状一致。你可以在心里做一次形状验算。如果x是(batch, in_features),W是(out_features, in_features),那么dz.T @ x的结果是(out_features, in_features),正好和W一致;dz @ W的结果是(batch, in_features),正好和x一致。
4.2 Linear Layer 反向代码示例
下面是一个简单例子,展示反向逻辑:
def linear_backward(dz, x, weight, bias): # dz: (batch, out_features) # x: (batch, in_features) # weight: (out_features, in_features) dw = dz.T @ x db = dz.sum(axis=0) dx = dz @ weight return dx, dw, db这段代码只适合二维输入。如果输入是三维(batch, seq_len, in_features),需要把batch * seq_len合并,或者在实现时用np.tensordot、torch.einsum处理。我推荐先写成二维形式,通过测试后再加维度扩展。
4.3 Embedding Layer 反向传播的要点
Embedding Layer 反向是新手最容易混淆的地方。前向是查表,反向则是把梯度放回被选中的那些行。
假设输入 token ID 是[2, 0, 1],前向输出了三行向量。反向传播时,每个输出向量会收到一个梯度向量,你需要把这些梯度向量累加到 embedding table 的第 2、0、1 行。如果有多个位置都选中了同一个 ID,梯度要相加。
一个朴素实现:
def embedding_backward(dout, token_ids, vocab_size, embedding_dim): # dout: (batch, seq_len, embedding_dim) dembedding = np.zeros((vocab_size, embedding_dim)) for b in range(dout.shape[0]): for s in range(dout.shape[1]): idx = token_ids[b, s] dembedding[idx] += dout[b, s] return dembedding这个实现性能不高,但逻辑清楚。实际做大规模训练时,会使用散列累加或 PyTorch 的自动微分实现。作业阶段,先保证逻辑正确,再考虑优化。
Embedding 层的梯度有一个特殊性:只有被输入选中的行梯度非零,其他行梯度保持 0。这意味着初始化时没有选中的行,在第一次迭代里不会更新,这是正常现象,不是 bug。
4.4 参数更新与梯度裁剪
拿到梯度后,参数更新通常用随机梯度下降最简单的版本:
weight -= learning_rate * dw bias -= learning_rate * db embedding_table -= learning_rate * dembedding这里有个小细节:如果手写的是多层的完整 Transformer,所有模块共用一个 loss,梯度需要从后往前逐层计算并累积。初学阶段可以在每个模块里只更新自己的参数,但要保证整个 forward 和 backward 使用的中间变量都保存在一个“缓存”结构里,否则反向传播时拿不到上一层输入。
梯度裁剪是训练稳定性的第一道防线。常见做法有两种:
- 按梯度范数裁剪:如果全局梯度范数超过阈值,就等比例缩小;
- 按梯度值裁剪:把每个梯度限制在
[-clip_value, clip_value]。
在 Transformer 训练中,我一般建议至少加上按范数裁剪,尤其当 batch size 较大、学习率偏高时。梯度裁剪虽然不会提升模型理论能力,但能有效避免个别异常样本把参数推出正常范围。
5. 训练稳定性:梯度检查、数值检查和常见报错排查
5.1 用梯度检查验证反向传播是否正确
写完整套反向传播后,最直接验证方法是数值梯度检查。原理很简单:对某个参数theta,给它加一个极小量epsilon和减一个极小量,分别计算 Loss,得到近似梯度:
grad_numerical = (loss(theta + eps) - loss(theta - eps)) / (2 * eps)再和你手写的反向梯度grad_analytic比较。如果两者差距小于 1e-5 左右,通常说明反向实现没问题。误差太大时检查两点:一是eps是否取太大,二是代码中是否有原地修改参数导致前向计算被污染。
这个检查一定要用很小模型做。如果直接拿完整 Transformer 检查,一个参数算两次前向,耗时很高,而且多个模块叠加后误差会被放大,不容易定位。
5.2 输出为 NaN、梯度爆炸、Embedding 梯度为 0 的排查顺序
我在实际跑手写模型时,最常遇到三类问题。
第一类是 Loss 直接变 NaN。优先按这个顺序查:
- 学习率是不是太大;
- 初始化标准差是不是太大;
- 有没有除零操作;
- 有没有在前向连续计算时出现中间变量被覆盖;
- 输入数据里是否存在异常大的值。
第二类是某个模块的梯度过大。看每一层梯度的范数,如果前面几层梯度远大于最后一层,说明反向传播过程中梯度累积过快。这时候先检查残差连接和 LayerNorm 是否已经实现,再检查初始化是否满足方差缩放。
第三类是 Embedding 梯度全为 0。先看输入 token ID 是不是越界,再看查表返回的梯度是否被赋值到正确位置,最后看是否在反向传播之前就清零了梯度缓冲区。我见过有人反复调用zero_grad(),结果把刚算出来的梯度也清掉了。
5.3 小批量训练时的经验边界
写完这套逻辑后,很多人会立刻把vocab_size设成 50000,开启完整训练。我不建议这么做。更稳妥的方式是先跑一个小批量,比如一个 batch 里只有 2 句话,每句话 8 个 token,训练 10 步,观察 Loss 是否从基数值缓慢下降。只要 Loss 在下降,反向传播和参数更新大概率是对的,再逐步增加 batch size 和序列长度。
小批量环境能暴露很多问题,但也会掩盖另一些问题。比如显存占用不大,不代表大 batch 下梯度范数稳定;序列短,不不代表长序列位置编码实现正确。所以小批量只是第一步,完整验证仍需要用小模型、中等序列长度跑一次。
注意:这里不要一上来就开最大并发或最长序列,先用单条样本确认整个 forward/backward 链路是通的,再扩展规模。
6. 从作业到后续扩展:多层 Transformer 中的参数共享与初始化策略
6.1 多层堆叠时的参数命名和统一管理
在单个 Linear Layer 反向传播正确之后,下一步就是把它放进一个多层结构里。建议把每层参数放进一个字典,比如:
params = { "embedding_table": ..., "transformer_block_0_linear1_w": ..., "transformer_block_0_linear1_b": ..., "transformer_block_0_linear2_w": ..., ... }统一管理的好处是:反向传播时可以按层逆序计算,更新时也能统一做梯度裁剪。如果每个模块的权重都散落在不同变量里,梯度检查和代码调试会变得很痛苦。
参数命名不是小事。后续加多头注意力、LayerNorm、残差连接时,没有清晰命名的代码会迅速失控。我一般会按“模块名 + 层级序号 + 参数类型”来命名。
6.2 残差、LayerNorm 与初始化尺度配合
当 Transformer 堆叠多层时,单看一层 Linear 的初始化可能合理,但整个网络输出方差会随层数增长。为了应对这一点,常见策略是在残差分支上做特殊缩放,例如把 Attention 和 FFN 里部分线性层的标准差乘以1 / sqrt(2 * num_layers)。
LayerNorm 也能稳定训练,但它不是万能的。LayerNorm 可以把每层输出重新归一化到合理尺度,但它不会自动修正错误的初始化选择。如果初始权重把前向信号压得太小,LayerNorm 之后依然可能让梯度消失。
我的建议是:先写一个 2 层的迷你 Transformer,用固定随机种子做梯度检查;确认通过后,再增加层数,观察前向方差。如果层数增加后 Loss 不降,优先怀疑初始化,而不是先调学习率。
6.3 后续扩展建议
如果你把 Linear Layer 和 Embedding Layer 都手写完了,并且反向传播通过梯度检查,接下来的扩展路径可以按这个顺序:
- 加入 LayerNorm 和残差连接;
- 加入多头注意力;
- 加入位置编码;
- 把 Embedding 和输出层共享权重;
- 加入学习率调度器;
- 加入 batch 级别的数据加载和日志记录。
每加一个模块,都先用同一个固定样例做梯度检查,不要等所有模块都写完才调。手写 Transformer 最怕的不是单模块错,而是多个模块叠加后,你根本不知道是哪一层出了问题。
回到开头说的:这节作业真正训练的是你对“矩阵运算在哪里、梯度从哪里来”的掌控感。参数初始化、Linear、Embedding、反向传播这些词,单独看都不难,但拼在一起时,很多错误是维度、尺度和累加位置造成的。我个人更建议先把单条任务跑稳,再考虑批量和完整训练。真正落地时,最该盯住的不是功能列表,而是输入形状、初始化尺度和梯度形状一致。踩过几次之后你会发现,大部分问题不是模型能力不够,而是最基础的模块没有做到可验证。