ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

手写Transformer核心组件:Linear与Embedding层的初始化及反向传播

手写Transformer核心组件:Linear与Embedding层的初始化及反向传播 最近在跟 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_outfan_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(axis0)这里面的维度规则是任何梯度的形状必须和原始变量形状一致。你可以在心里做一次形状验算。如果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(axis0) 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、反向传播这些词单独看都不难但拼在一起时很多错误是维度、尺度和累加位置造成的。我个人更建议先把单条任务跑稳再考虑批量和完整训练。真正落地时最该盯住的不是功能列表而是输入形状、初始化尺度和梯度形状一致。踩过几次之后你会发现大部分问题不是模型能力不够而是最基础的模块没有做到可验证。
返回列表