Skip to content

归一化与残差连接 ​

标签
AI/llm/架构
字数
4169 字
阅读时间
17 分钟

11-Self-Attention 机制 和 12-FFN 与激活函数 是两个子模块本身。这一篇讲它们之间怎么连 —— 残差连接解决梯度能不能传回去,归一化解决各层输入分布漂不漂。

两者合起来,是「为什么几十上百层的网络还能训起来」的工程答案。

要解决的两个问题 ​

梯度消失与爆炸 ​

反向传播是链式法则的递归应用。L 层网络的梯度是一连串偏导的乘法:

∂Loss∂W1=∂Loss∂fL⋅∂fL∂fL−1⋯∂f2∂f1⋅∂f1∂W1

设每层雅可比矩阵的谱范数为 σ,连乘 L 层后:

σ结果
<1梯度以 σL 指数衰减 —— 深层参数收不到有意义的信号
>1梯度以 σL 指数增长 —— 更新剧烈震荡甚至 NaN
=1才稳定传播,但实际训练中几乎不可能自然满足

退化问题 ​

2015 年何恺明等人发现一个反直觉现象:更深的网络反而比浅层网络表现更差,而且不是过拟合(训练误差也更高),是优化本身出了问题。

理论上 56 层不该比 20 层差 —— 多出的 36 层只要学成恒等映射就能持平。但让网络「学会什么都不做」远比想象的困难。 这个观察直接催生了残差连接。

残差连接 ​

output=x+F(x)

F(x) 叫「残差函数」。ResNet 用这一个加号把网络从几十层推到 152 层乃至上千层。

为什么学残差更容易 ​

两个角度:

认知负担角度。 无残差 = 从零画一份终版图纸;有残差 = 在初版图纸上标修改意见,不用改的部分自动保留。当某层确实不需要做什么时,F(x) 只需趋近于零 —— 而学一个趋近零的函数远比学恒等映射容易。初始化时权重通常接近零小值,F(x) 天然接近零,网络自动具备「什么都不做」的能力作为起点。

数学角度。 展开从第 L 层回传到第 l 层的梯度:

∂L∂xl=∂L∂xL∏k=lL−1(I+∂Fk∂xk)

把这个连乘展开后,必然包含一项恒为 I 的乘积(所有括号都取 I 那一项):

∂L∂xl=∂L∂xL⋅(I+其他交叉项)

无论中间各层的 ∂Fl/∂xl 有多小甚至趋近于零,梯度都能通过这条「恒等通道」原封不动地传回去。

这就是「梯度高速公路」的确切含义 —— 它在反向计算图里开辟了一条短路,梯度不必走那条经过所有非线性变换的崎岖小路。纯乘法链变成了「加法 + 乘法」的混合结构:纯乘法链对数值极其敏感(连续乘 0.9 五十次试试),加法让梯度有了保底值。

Transformer 里每个 Block 有两条残差:

h=x+Attention(x)output=h+FFN(h)

一个 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 独立归一化:

μ=1d∑j=1dxj,σ2=1d∑j=1d(xj−μ)2,yj=γjxj−μσ2+ϵ+βj

四个特性正好一一对应 BatchNorm 的四个问题:

  1. 完全独立于 batch —— 只依赖自身的 d 维特征
  2. 不受序列长度影响 —— 各 token 独立归一化
  3. 训练与推理行为完全一致 —— 没有 running stats
  4. 自回归友好 —— 每个 token 可独立归一化

除以 $d$ 不是 $d-1$

这里用的是总体方差形式。统计库的默认 var() 常是无偏形式(除以 n−1),两者在小 d 下差异可观。复现算子时要查 API 语义 —— 见 04-概率、Softmax 与信息论。

RMSNorm ​

RMS(x)=1d∑jxj2,yj=γjxjRMS(x)

核心假设(Zhang & Sennrich, 2019):LayerNorm 的效果主要来自「重新缩放」而非「重新中心化」 —— 减去均值那一步或许不是必需的。所以 RMSNorm 去掉均值计算,并且通常去掉 β(没有 re-centering 就不需要偏移项)。

两者关系的数学依据:

RMS(x)2=mean(x2)=Var(x)+mean(x)2

当 mean(x) 接近 0 时,RMS(x) 近似等于 std(x) —— 而深层 Transformer 中特征向量的均值通常确实接近零(尤其在 Pre-Norm 架构下,残差不断累加,均值的相对比重越来越小)。这为 RMSNorm 的有效性提供了直觉解释。

归一化的本质作用是控制激活值的尺度,防止某些维度数值过大或过小导致梯度不稳。这个目标靠除以 RMS 已基本达成;「去除直流分量」在深度网络里不是必需品 —— 每层的可学习参数已足以隐式处理均值偏移。

计算量优势是真的,而且恰好在瓶颈上:

需要几次扫描
LayerNorm实际实现通常两遍:第一遍同时算 ∑x 与 ∑x2,第二遍归一化
RMSNorm一遍算 ∑x2,一遍归一化

这类操作是 memory-bound 的(瓶颈在显存带宽而非算力),少一次全局归约扫掠就是少一次 HBM 读取 —— 收益直接体现在墙上时间里。机制见 01-GPU 硬件架构与存储层次 的 Roofline 一节。

三者对照 ​

维度BatchNormLayerNormRMSNorm
归一化维度batch 维(跨样本)特征维(每样本独立)特征维(每样本独立)
去均值是是否
可学习参数γ+βγ+β仅 γ
依赖 batch强,小 batch 不稳无无
训练/推理一致不一致(需 running stats)一致一致
计算量高(需跨 batch 通信)中(两遍扫描)较低(可一遍)
适用CV(CNN + 大 batch)NLP / Transformer大规模 LLM
代表ResNet、EfficientNetGPT-2/3、BERTLLaMA、Mistral、Qwen

三种归一化的差别在「沿哪个维度统计」。把张量形状记作 (N,L,d)——N 个样本、每个 L 个 token、每个 token d 维:

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 的雅可比矩阵调制一次:

∂L∂xl=∂L∂xL∏k=lL−1[∂LNk∂(xk+Fk)⋅(I+∂Fk∂xk)]

LayerNorm 的雅可比不是单位阵 —— 它涉及均值与方差的梯度,会对回传梯度做一次非平凡变换。这些变换累乘起来可能显著改变梯度的方向与大小。96 层的 GPT-3 规模下,累积效应可能引发爆炸,必须靠学习率 warmup 抑制。

Pre-Norm 的残差加法之后没有任何操作:

∂L∂xl=∂L∂xL∏k=lL−1[I+∂(Fk(LNk(xk)))∂xk]

每个因子仍是「I + 某个东西」的形式 —— 恒等路径直达任意浅层,LayerNorm 的雅可比只出现在交叉项里,不阻塞主梯度通路。

一句话: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:想兼得的混合方案 ​

output=LayerNorm(α⋅x+SubLayer(x))

α>1 放大残差路径贡献,同时 SubLayer 参数按规则缩小初始化(β<1),α 与 β 由深度 L 决定。形式上仍是 Post-Norm(可能拿到更好的最终效果),但通过放大残差路径实现了类似 Pre-Norm 的梯度稳定性。微软用 DeepNorm 训到了 1000 层。

不过在实际开源大模型里采用率远不如 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 不大(d=4096 的 FP16 约 8 KB),但 LayerNorm 在每个 Block 被调用两次(Attention 前、FFN 前),乘上几十个 Block、batch 与序列长度,累计访存量相当可观。

优化方向两条:

  • Welford 在线算法 —— 一遍同时算均值与方差,同时避开数值不稳定的 1N∑xi2−(1N∑xi)2 形式。代价是累加本身是顺序的,要在并行归约框架下仔细实现。见 01-数值计算与精度
  • RMSNorm 的天然优势 —— 去掉均值计算后只需要一个归约(∑x2)而非两个,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 往返:它还省掉了一份中间张量的显存。在几十层堆叠的模型里,这类「逐元素 + 归约」的融合是最容易拿到确定性收益的一类优化。

相关 ​

参考 ​

贡献者 ​

文件历史 ​