ARTICLE DETAIL

资讯详情

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

手撕 Transformer Block:60 行 PyTorch 跑通 FFN、残差连接与 LayerNorm

手撕 Transformer Block:60 行 PyTorch 跑通 FFN、残差连接与 LayerNorm 手撕 Transformer Block60 行 PyTorch 跑通 FFN、残差连接与 LayerNorm前两篇写了大模型 API 省 token 的技巧和最小 ReAct Agent这篇往下沉一层模型内部一个 Transformer Block 到底怎么工作。收藏榜上讲 FFN/残差/LayerNorm 概念的帖子正在收收藏但多数是纯概念文——本篇直接手撕60 行 PyTorch 实现完整 Block带不花算力的数值自检。结论先放这儿组件干什么一句话铁律FFN对 Attention 结果做非线性加工2/3 的参数在这调参先看它残差连接让网络只学“变化量”梯度的高速公路深网训得动的前提LayerNorm稳定隐藏表示的分布Pre-LN 在前Post-LN 在后深了用 Pre-LN数据流一句话Attention 管“信息从哪来”FFN 管“拿到信息怎么加工”残差管“只学变化量”LayerNorm 管“数值稳定”。一、完整代码单文件直接跑# transformer_block.py — 标准 Pre-LN Block # 依赖pip install torch import torch import torch.nn as nn class TransformerBlock(nn.Module): def __init__(self, d_model512, n_heads8, d_ff2048, dropout0.1): super().__init__() self.ln1 nn.LayerNorm(d_model) # Pre-LNnorm 在子层前面 self.attn nn.MultiheadAttention( d_model, n_heads, dropoutdropout, batch_firstTrue) self.ln2 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), # 现代实现用 GELU不是 ReLU nn.Linear(d_ff, d_model), nn.Dropout(dropout), ) def forward(self, x, pad_maskNone): # ---- 多头注意力 残差 ---- h self.ln1(x) a, _ self.attn(h, h, h, key_padding_maskpad_mask, need_weightsFalse) x x a # 残差学的是变化量 # ---- FFN 残差 ---- x x self.ffn(self.ln2(x)) return x不花算力的数值自检验证 shape、数值稳定、梯度流通、FFN 参数占比if __name__ __main__: torch.manual_seed(0) blk TransformerBlock(d_model64, n_heads4, d_ff256).eval() x torch.randn(2, 10, 64, requires_gradTrue) # (batch, seq, d_model) out blk(x) assert out.shape x.shape # shape 不变是堆叠多层的前提 assert torch.isfinite(out).all() # 数值不炸 out.sum().backward() assert x.grad is not None and torch.isfinite(x.grad).all() # 残差在梯度能从输出直达输入几十层也训得动 n_ffn sum(p.numel() for p in blk.ffn.parameters()) n_all sum(p.numel() for p in blk.parameters()) print(fFFN 参数占比: {n_ffn / n_all:.0%}) # 实测约 66% print(self-check ok)跑一下四条 assert 全过FFN 参数占比打印约 66%——铁律“2/3 的参数在 FFN”当场验证。二、FFN参数大头不是配角Attention 常被当主角但看参数量两个d_model × d_ff矩阵占了整个 Block 约 2/3 的参数和计算量。这就是为什么中间维度习惯取 4 倍512 → 2048——给非线性加工留容量。为什么一定要非线性没有 GELU/ReLU两层线性叠加还是线性Block 的加工能力直接缩水成一层。三、残差连接网络只学“变化量”x x attn(...)这一行是整个深度学习的经验结晶网络不学完整表示只学相对输入的增量。两个直接后果梯度直达反向传播时梯度可以沿x f(x)的恒等分支一路传到第一层几十层的网络训得动上面自检里 x.grad 的 assert 就是在验这个好训每层只需要拟合“修正量”初始化时 Block 近似恒等映射深堆不崩四、LayerNormPre-LN 还是 Post-LN差别就一行代码的位置# Post-LN原论文norm 在残差之后 x self.ln1(x self.attn(x, x, x)[0]) # Pre-LNGPT/Llama 等现代实现norm 在子层之前 x x self.attn(self.ln1(x), self.ln1(x), self.ln1(x))[0]Post-LN 深了难训梯度会被 norm 反复缩放必须小心 warmupPre-LN 的残差分支是干净的恒等路径深堆稳定代价是最后一层外面要补一个 final LayerNorm。新写代码默认 Pre-LN。另外别丢epsLayerNorm 的 eps 防 0 方差除零量化/半精度部署时把 eps 设成 0 会直接 NaN。五、三个踩坑写 Block 时都遇到过Pre/Post-LN 混着堆一半层 Pre 一半层 Post训练曲线直接崩。选定一种全层统一GELU 写成 ReLU复现 BERT/GPT 系结构时激活函数对不上loss 曲线有差异但不报错很难查need_weightsTrue忘关MultiheadAttention 默认返回注意力权重矩阵batch 大时白吃一份显存推理慢一截总结一个 Block Attention信息交互 FFN信息加工 残差只学变化量 LayerNorm数值稳定。铁律压成三句2/3 的参数在 FFN调参先看它残差是梯度高速公路没有它就没有深堆新代码默认 Pre-LNeps 别丢下一篇可以接着往下走Decoder 怎么用这个 Block 逐 token 生成文本因果掩码 KV Cache感兴趣的点个收藏。
返回列表