归一化与残差连接
11-Self-Attention 机制 和 12-FFN 与激活函数 是两个子模块本身。这一篇讲它们之间怎么连 —— 残差连接解决梯度能不能传回去,归一化解决各层输入分布漂不漂。
两者合起来,是「为什么几十上百层的网络还能训起来」的工程答案。
要解决的两个问题
梯度消失与爆炸
反向传播是链式法则的递归应用。
设每层雅可比矩阵的谱范数为
| 结果 | |
|---|---|
| 梯度以 | |
| 梯度以 | |
| 才稳定传播,但实际训练中几乎不可能自然满足 |
退化问题
2015 年何恺明等人发现一个反直觉现象:更深的网络反而比浅层网络表现更差,而且不是过拟合(训练误差也更高),是优化本身出了问题。
理论上 56 层不该比 20 层差 —— 多出的 36 层只要学成恒等映射就能持平。但让网络「学会什么都不做」远比想象的困难。 这个观察直接催生了残差连接。
残差连接
为什么学残差更容易
两个角度:
认知负担角度。 无残差 = 从零画一份终版图纸;有残差 = 在初版图纸上标修改意见,不用改的部分自动保留。当某层确实不需要做什么时,
数学角度。 展开从第
把这个连乘展开后,必然包含一项恒为
无论中间各层的
有多小甚至趋近于零,梯度都能通过这条「恒等通道」原封不动地传回去。 这就是「梯度高速公路」的确切含义 —— 它在反向计算图里开辟了一条短路,梯度不必走那条经过所有非线性变换的崎岖小路。纯乘法链变成了「加法 + 乘法」的混合结构:纯乘法链对数值极其敏感(连续乘 0.9 五十次试试),加法让梯度有了保底值。
Transformer 里每个 Block 有两条残差:
一个 32 Block 的模型,信号从第 1 层到第 32 层要经过 64 个子层 —— 没有残差,梯度几乎不可能穿越。
残差的反向通道里有一条恒为 I 的路径。
无残差:L 层就是 L 个雅可比矩阵连乘,每个因子的谱范数 σ 决定结局
∂L/∂x₁ = ∂L/∂x_L · J_L · J_{L−1} · … · J₁
σ < 1 ──▶ 梯度以 σ^L 指数衰减,深层参数收不到有意义的信号
σ > 1 ──▶ 指数增长,更新震荡甚至 NaN
有残差:每个因子变成 (I + ∂F_k/∂x_k)
∂L/∂x_l = ∂L/∂x_L · ∏_k (I + ∂F_k/∂x_k)
└─ 把这个连乘展开,必然包含「每个括号都取 I」的那一项
那一项恒等于单位阵,与各层的权重无关
正向 x ─────────────────────────▶ x + F(x) 残差路径上只有加法
└──▶ F(x) ──▶ ┘
反向 ∂L/∂x_L ══════════════════▶ ∂L/∂x_l 恒等通道,不经过任何非线性
└─ 其余交叉项走 F 那条路 ─┘
中间各层的 ∂F/∂x 无论多小、甚至趋近于零,梯度都能沿恒等通道原封不动传回去。
Transformer 里每个 Block 有两条残差(Attention 一条、FFN 一条),
32 层就是 64 个子层 —— 没有这条通道,梯度几乎不可能穿越。归一化全景
残差解决「梯度能不能传回去」,归一化解决另一个问题:特征分布漂移 —— 随着训练进行,每层输入分布不断变化,后层必须不断适应前层的「新脾气」,训练不稳定且收敛慢。
归一化的目标是把每层输入拉回稳定分布区间。
BatchNorm 为什么不适合 NLP
BatchNorm 沿 batch 维度统计。它在 CV 上很成功,但搬到 Transformer 有四个根本困难:
| 问题 | 说明 |
|---|---|
| 序列长度不等 | padding 位置的特征值无意义,纳入均值/方差会引入噪声 |
| Batch 统计不稳定 | 大模型单卡 micro batch 可能只有 1–4,用它估计统计量方差极大 |
| 训练/推理不一致 | 训练用 batch 统计量,推理用 running mean/variance,两个分布未必匹配 |
| 自回归不友好 | 逐 token 生成时 batch size 为 1,batch 统计量完全没有意义 |
LayerNorm
改成沿特征维度、对每个 token 独立归一化:
四个特性正好一一对应 BatchNorm 的四个问题:
- 完全独立于 batch —— 只依赖自身的
维特征 - 不受序列长度影响 —— 各 token 独立归一化
- 训练与推理行为完全一致 —— 没有 running stats
- 自回归友好 —— 每个 token 可独立归一化
除以 $d$ 不是 $d-1$
这里用的是总体方差形式。统计库的默认 var() 常是无偏形式(除以
RMSNorm
核心假设(Zhang & Sennrich, 2019):LayerNorm 的效果主要来自「重新缩放」而非「重新中心化」 —— 减去均值那一步或许不是必需的。所以 RMSNorm 去掉均值计算,并且通常去掉
两者关系的数学依据:
当
归一化的本质作用是控制激活值的尺度,防止某些维度数值过大或过小导致梯度不稳。这个目标靠除以 RMS 已基本达成;「去除直流分量」在深度网络里不是必需品 —— 每层的可学习参数已足以隐式处理均值偏移。
计算量优势是真的,而且恰好在瓶颈上:
| 需要几次扫描 | |
|---|---|
| LayerNorm | 实际实现通常两遍:第一遍同时算 |
| RMSNorm | 一遍算 |
这类操作是 memory-bound 的(瓶颈在显存带宽而非算力),少一次全局归约扫掠就是少一次 HBM 读取 —— 收益直接体现在墙上时间里。机制见 01-GPU 硬件架构与存储层次 的 Roofline 一节。
三者对照
| 维度 | BatchNorm | LayerNorm | RMSNorm |
|---|---|---|---|
| 归一化维度 | batch 维(跨样本) | 特征维(每样本独立) | 特征维(每样本独立) |
| 去均值 | 是 | 是 | 否 |
| 可学习参数 | 仅 | ||
| 依赖 batch | 强,小 batch 不稳 | 无 | 无 |
| 训练/推理一致 | 不一致(需 running stats) | 一致 | 一致 |
| 计算量 | 高(需跨 batch 通信) | 中(两遍扫描) | 较低(可一遍) |
| 适用 | CV(CNN + 大 batch) | NLP / Transformer | 大规模 LLM |
| 代表 | ResNet、EfficientNet | GPT-2/3、BERT | LLaMA、Mistral、Qwen |
三种归一化的差别在「沿哪个维度统计」。把张量形状记作
BatchNorm:沿 N(与 L)统计,每个特征维 d 单独算一组 mean / variance
└─ 统计量依赖 batch 的内容
└─ padding 位置的特征值无意义,纳入统计会引入噪声
└─ 大模型单卡 micro batch 可能只有 1–4,用它估统计量方差极大
└─ 推理时 batch 为 1,统计量完全没有意义 ──▶ 自回归不友好
LayerNorm:沿 d 统计,每个 (样本, token) 独立算一组
│ N 与 L 两个维度都不参与
├─ 完全独立于 batch
├─ 不受序列长度影响
├─ 训练与推理行为一致(没有 running stats)
└─ 自回归友好:每个 token 可以独立归一化
RMSNorm:同样沿 d 统计,但只算 √mean(x²)
├─ 不减均值(不做 re-centering),通常也去掉 β
├─ 依据:RMS² = Var(x) + mean(x)²,深层 Transformer 的特征均值接近 0,
│ 此时 RMS 近似等于 std —— 减均值那一步不是必需的
└─ 收益:归约从两个(Σx 与 Σx²)减到一个(Σx²)Pre-Norm vs Post-Norm
LayerNorm 放在子层前还是后,看起来只是顺序微调,但对大模型的训练稳定性有决定性影响。
Post-Norm(原始 Transformer) Pre-Norm(GPT-2 起,现主流)
x x
│ │
├──────────────┐ ├──────────────┐
│ ▼ │ ▼
│ SubLayer(x) │ LayerNorm(x) ──▶ SubLayer
│ │ │ │
├──────────────┘ ← 残差加法 ├──────────────┘ ← 残差加法
▼ ▼
LayerNorm(x + SubLayer(x)) x + SubLayer(LN(x))
│ │
└─ 残差路径的终点还要过一层 LN └─ 残差路径到此结束,没有任何后续操作
回传时每过一个 Block,都被它的
雅可比矩阵调制一次
一句话差别:Pre-Norm 让残差路径保持「纯净」—— LayerNorm 的雅可比只出现在交叉项里。差别只在「残差路径是否干净」
Post-Norm 的残差加法结果要再经过一次 LayerNorm。回传时每经过一个 Block 都被 LayerNorm 的雅可比矩阵调制一次:
LayerNorm 的雅可比不是单位阵 —— 它涉及均值与方差的梯度,会对回传梯度做一次非平凡变换。这些变换累乘起来可能显著改变梯度的方向与大小。96 层的 GPT-3 规模下,累积效应可能引发爆炸,必须靠学习率 warmup 抑制。
Pre-Norm 的残差加法之后没有任何操作:
每个因子仍是「
一句话:Pre-Norm 让残差路径保持「纯净」。
实验证据
Xiong et al.(2020, On Layer Normalization in the Transformer Architecture)的理论分析指出:Post-Norm 在初始化时各层输出的方差会随深度线性增长,必须用较小学习率和较长 warmup 抑制;Pre-Norm 各层方差天然稳定,可以省略 warmup 直接用较大学习率。
另外,有研究表明 Post-Norm 在训练充分收敛时效果略好于 Pre-Norm(约 0.1–0.5 BLEU 点),但这需要极精细的超参调优。模型规模到数十亿参数以上时,Post-Norm 经常出现 loss spike 甚至训练发散,而 Pre-Norm 稳定。
训练一个 175B 模型可能花数百万美元算力。一次发散意味着前功尽弃。 在这种约束下,Pre-Norm 提供的稳定性保障远比那 0.1–0.5 BLEU 更有工程价值 —— 这就是它成为默认选择的原因,不是因为它效果更好。
从 GPT-2(2019)开始,GPT-3、PaLM、LLaMA、Mistral、Qwen 几乎全部采用 Pre-Norm。
DeepNorm:想兼得的混合方案
不过在实际开源大模型里采用率远不如 Pre-Norm —— Pre-Norm 足够好用、实现简单、社区验证充分,没有引入额外超参数的必要。
Kernel 视角
LayerNorm 是一个 reduction
数学不复杂,但在 GPU 上实现得讲究效率:核心挑战在于它需要对整个特征维做归约,而 GPU 擅长的是大规模并行的逐元素计算。
朴素实现需要两遍扫描:
朴素实现:两遍扫描,特征向量从 HBM 读两次
第一遍 第二遍
HBM ──▶ 片上 HBM ──▶ 片上
│ 读 x(d 个元素) │ 再读一遍 x
│ 算 sum(x) 与 sum(x²) │ 用 mean / variance 做归一化
│ │
└─▶ mean、variance └─▶ 输出 y
一个 token 不大(d = 4096 的 FP16 约 8 KB),但 LayerNorm 在每个 Block
被调用两次(Attention 前、FFN 前),再乘上几十个 Block、batch 与序列长度 ——
累计访存量相当可观。这类操作是 memory-bound 的,少扫一遍就是少一次 HBM 读取。特征向量的数据要从 HBM 读两次。单看一个 token 不大(
优化方向两条:
- Welford 在线算法 —— 一遍同时算均值与方差,同时避开数值不稳定的
形式。代价是累加本身是顺序的,要在并行归约框架下仔细实现。见 01-数值计算与精度 - RMSNorm 的天然优势 —— 去掉均值计算后只需要一个归约(
)而非两个,kernel 更简单,数值稳定性问题也更少。计算量减少只是附带收益,归约数量减半才是主要收益
Residual Add + LayerNorm 融合
Pre-Norm 架构下两者数据流上紧密相连:
未融合:加法结果先写回 HBM,LayerNorm 再读回来
┌────────────────┐
│ Residual Add │ 读 x、读 SubLayer 输出
└───────┬────────┘
│ 写 h ──▶ HBM
▼
┌────────────────┐
│ LayerNorm │ 从 HBM 读回 h
└───────┬────────┘
│ 写 next_input ──▶ HBM
▼
中间张量 h 在 HBM 上往返一次
融合:加法结果留在片上直接参与归约
┌───────────────────────────────────────────┐
│ Residual Add + LayerNorm 合成一个 kernel │
│ 读 x、读 SubLayer 输出 │
│ 加法结果留在片上 ──▶ 直接归约 ──▶ 归一化 │
└──────────────────┬────────────────────────┘
│ 写 next_input ──▶ HBM
▼
省掉一次 HBM 往返,同时省掉一份中间张量的显存既然 LayerNorm 必须读一遍 h,而 h 刚刚由加法产生,就没有理由先把 h 写回 HBM 再读回来 —— 把它们融合成一个 kernel,加法结果留在片上直接参与归约。这是「中间张量能否融合消除」这条检查表条目(见 03-线性代数与 GEMM)在真实模型里的标准应用。
融合的价值不只省一次 HBM 往返:它还省掉了一份中间张量的显存。在几十层堆叠的模型里,这类「逐元素 + 归约」的融合是最容易拿到确定性收益的一类优化。
相关
- 11-Self-Attention 机制 —— 被残差包裹的第一个子层
- 12-FFN 与激活函数 —— 第二个子层
- 02-Transformer 架构 —— 整体骨架
- 05-反向传播与梯度优化 —— 梯度消失爆炸的一般性讨论
- 01-数值计算与精度 —— 方差的稳定算法与 Welford
- 01-GPU 硬件架构与存储层次 —— 为什么 memory-bound 操作少扫一遍就有效
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块一-前置知识/transformer/36-layernorm与残差连接深入理解
- https://arxiv.org/abs/1512.03385
- https://arxiv.org/abs/1607.06450
- https://arxiv.org/abs/1910.07467
- https://arxiv.org/abs/2002.04745
- https://arxiv.org/abs/2203.00555
YJ