一篇读懂归一化:LayerNorm、RMSNorm 与 Pre-Norm 的训练稳定性账

一篇读懂归一化:LayerNorm、RMSNorm 与 Pre-Norm 的训练稳定性账

它解决什么问题

深层网络训练有两个老绊脚石。其一,各层输入的尺度漂移:前面的参数一更新,后面层看到的激活分布就变,训练像在追一块移动的地板——BatchNorm 论文把它归因为「内部协变量偏移」。其二,梯度不稳定:反向传播要穿过 L 层,尺度信号被反复放大或缩小,要么爆炸要么消失。

BatchNorm(2015)的解法是用 mini-batch 的均值和方差把每层输入拉回稳定尺度。但它的统计量来自 batch,这在三类场景里水土不服:batch 很小时方差估计噪声大;序列任务里样本长短不一,RNN 每个时间步还得单独算一套统计量;训练用 batch 统计、推理用滑动平均,两套计算不一致。

LayerNorm 论文(Ba 等,2016)点名了这三条,然后换了统计量的来源:不看 batch,只看当前样本在这一层内的全部特征。这一个改动,让它后来统治了整个 Transformer 时代。

它怎么工作

LayerNorm:把统计量搬回单样本

对一层内 n 个特征:

μ = (1/n)·Σᵢ xᵢ,σ² = (1/n)·Σᵢ (xᵢ − μ)²

y = γ ⊙ (x − μ)/√(σ² + ε) + β

γ(增益)与 β(偏移)是逐特征的可学习参数,ε 防止除零。训练与推理执行完全相同的计算,不需要缓存任何统计量;论文实验显示它能稳定 RNN 隐状态动态、大幅缩短训练时间。

注意公式里藏着两个不变性:减均值带来重中心化不变性(整体平移输入,输出不变),除以标准差带来重缩放不变性(整体缩放输入,输出不变)。重缩放不变性意味着激活尺度的变化不再影响该层的有效增益——模型自己获得了隐式的学习率自适应。

RMSNorm:减法是多余的

Zhang 与 Sennrich(2019)提出一个大胆的假设:LayerNorm 的成功靠的是重缩放不变性,重中心化可以省掉。于是:

RMS(x) = √((1/n)·Σᵢ xᵢ²),y = g ⊙ x/RMS(x)

不中心化,就省掉一次均值计算和一半的统计量;β 也没有了,只剩增益 g。这篇论文的效率消融在单卡上实测了每千步训练时间:Transformer base 从 248s 降到 231s(省 6.9%),RNNSearch 在 Theano 上从 988s 降到 652s(省 34.0%),order-embedding 模型从 12.02s 降到 7.12s(省 40.8%,用只采样 6.25% 特征估计 RMS 的 pRMSNorm 可到 63.9%)。论文总结:不同模型上省 7%~64% 运行时间,性能与 LayerNorm 相当。

效果如何?截至 2026-10-07,LLaMA、Qwen、DeepSeek 等主流开源大模型全线使用 RMSNorm;LLaMA 风格的实现连加性偏置也一并省去,每层只剩一个增益向量 g。省下的每个 FLOPs,乘上万亿 token 的训练量都是真金白银。

用 numpy 验证两件事

先验证 LayerNorm 与 RMSNorm 的数值关系:均值近零时两者应当几乎重合,均值偏移大时应当分开;再验证重缩放不变性。

import numpy as np

def layer_norm(x, g, b, eps=1e-5):
    mu = x.mean(axis=-1, keepdims=True)
    var = x.var(axis=-1, keepdims=True)
    return g * (x - mu) / np.sqrt(var + eps) + b

def rms_norm(x, g, eps=1e-5):
    rms = np.sqrt(np.mean(x * x, axis=-1, keepdims=True) + eps)
    return g * x / rms

rng = np.random.default_rng(42)
g, b = np.ones(512), np.zeros(512)
x0 = rng.standard_normal(512)
x1 = x0 + 10.0
for name, x in [("均值≈0 ", x0), ("均值≈10", x1)]:
    diff = np.max(np.abs(layer_norm(x, g, b) - rms_norm(x, g)))
    print(f"{name}: LayerNorm 与 RMSNorm 输出最大差 = {diff:.4f}")
print(f"RMSNorm 重缩放不变性(x→7x): {np.max(np.abs(rms_norm(7 * x0, g) - rms_norm(x0, g))):.2e}")

运行输出:

均值≈0 : LayerNorm 与 RMSNorm 输出最大差 = 0.0189
均值≈10: LayerNorm 与 RMSNorm 输出最大差 = 3.7342
RMSNorm 重缩放不变性(x→7x): 1.58e-05

标准正态激活的均值约 −0.018,两种归一化的输出最大差只有 0.019——这就是「中心化在零均值附近几乎不起作用」的直观形态;把输入整体抬高 10 之后差异放大到 3.7,说明省掉中心化不是无条件的免费午餐(见「边界」)。重缩放不变性的残差 1.6e-05 恰好是 ε 的量级。

Pre-Norm 与 Post-Norm:归一化挂在哪,梯度账就怎么算

Transformer 里还有一个自由度:LN 挂在残差的哪一侧。

  • Post-Norm(原始 Transformer,2017):h_{l+1} = LN(h_l + Sublayer(h_l))——归一化在残差求和之后;
  • Pre-Norm(GPT-2 起的主流):h_{l+1} = h_l + Sublayer(LN(h_l))——归一化在残差分支内部。

差别只是位置,但反向传播的结构完全不同。LN 的雅可比是 (1/σ)·(I − 11ᵀ/n − x̂x̂ᵀ/n):要除以前向标准差,还要投影掉常数方向和 x̂ 方向。Post-Norm 下,主干梯度的每一跳都要穿过这个雅可比,等于每层交一次「收缩税」;Pre-Norm 下主干是恒等映射,梯度有一条不衰减的直达通道,LN 的税只落在分支贡献上。

Xiong 等(ICML 2020)把这笔账算清了。Post-LN:输出端的参数梯度范数量级为 O(d·√(ln d)),与深度无关地大;而附录 F 的初步分析显示,梯度传到第 l 层要带上约 (2/3)^((L−l)/2) 的衰减因子——越往输入端越弱。初始化时输出端梯度巨大,学习率稍大第一步就发散,这就是 Post-LN 必须配学习率 warmup 的原因:论文实测 IWSLT14 翻译上,Post-LN 不加 warmup 的 BLEU 只有 8.45,加了 4000 步 warmup 后约 34。Pre-LN:残差主干直通使隐状态范数随深度线性增长(论文引理 2),顶层的参数梯度被 1/√L 重标定为 O(d·√(ln d/L)),各层量级均衡——warmup 可以整个去掉,BERT 预训练实测提速约 40%。

把梯度账跑出来

理论之外,我们用一个放大率模型把账数值化:把每个子层近似为两层 ReLU FFN(Xavier 尺度初始化,γ=1、β=0),前向记录隐状态范数,反向从顶层单位梯度手推雅可比回传,重复 4 个随机种子取平均。以下数值均由这段代码实际运行得到。

def ln_forward(x, eps=1e-5):
    return (x - x.mean()) / np.sqrt(x.var() + eps)

def ln_jacobian(x, eps=1e-5):
    """LayerNorm 雅可比 J = (1/s)[I − 11ᵀ/d − x̂x̂ᵀ/d],x̂ 为标准化后的 x。"""
    d = x.shape[-1]
    s = np.sqrt(x.var() + eps)
    xh = (x - x.mean()) / s
    return (np.eye(d) - np.ones((d, d)) / d - np.outer(xh, xh) / d) / s

def simulate(norm, L, d=128, n_seed=4):
    """把每个子层近似为两层 ReLU FFN(Xavier 尺度),数值反传,记录三类范数。"""
    k = 4 * d
    hs = np.zeros((n_seed, L))   # 前向隐状态范数 /√d
    dh = np.zeros((n_seed, L))   # 反向隐状态梯度范数
    dp = np.zeros((n_seed, L))   # 反向 FFN 第二层权重的参数梯度范数
    for si in range(n_seed):
        r = np.random.default_rng(si)
        h = r.standard_normal(d)
        Ws, cache = [], []
        for _ in range(L):                                  # 前向
            W1 = r.standard_normal((k, d)) / np.sqrt(d)
            W2 = r.standard_normal((d, k)) / np.sqrt(k)
            a1 = W1 @ (ln_forward(h) if norm == "pre" else h)
            f = W2 @ np.maximum(a1, 0)
            cache.append((h.copy(), a1))
            h = h + f if norm == "pre" else ln_forward(h + f)
            Ws.append((W1, W2))
            hs[si, len(Ws) - 1] = np.linalg.norm(h) / np.sqrt(d)
        gvec = r.standard_normal(d)
        gvec /= np.linalg.norm(gvec)                        # 顶层单位隐状态梯度
        for l in range(L - 1, -1, -1):                      # 反向
            W1, W2 = Ws[l]
            h_in, a1 = cache[l]
            mask = (a1 > 0).astype(float)
            if norm == "pre":      # 分支挂在恒等主干上:δ 直传
                gz = gvec
                gvec = gvec + ln_jacobian(h_in).T @ (W1.T @ (mask * (W2.T @ gz)))
            else:                  # 分支求和之后要过 LN:δ = Jᵀ·g
                z = h_in + W2 @ np.maximum(a1, 0)
                gz = ln_jacobian(z).T @ gvec
                gvec = gz + W1.T @ (mask * (W2.T @ gz))
            dh[si, l] = np.linalg.norm(gz)
            dp[si, l] = np.linalg.norm(gz) * np.linalg.norm(np.maximum(a1, 0))
    return hs.mean(axis=0), dh.mean(axis=0), dp.mean(axis=0)

for norm in ["post", "pre"]:
    hs, dh, dp = simulate(norm, L=64)
    pick = [0, 16, 31, 47, 63]
    print(f"\n{norm}-Norm, L=64:  前向范数/√d " +
          ", ".join(f"{hs[l]:.2f}" for l in pick) +
          " | 隐状态梯度 " + ", ".join(f"{dh[l]:.2f}" for l in pick) +
          " | 参数梯度 " + ", ".join(f"{dp[l]:.1f}" for l in pick))
    print(f"  输入端/输出端参数梯度比 = {dp[0] / dp[-1]:.3f}")

print("\n深度敏感性(输入端/输出端参数梯度比):")
for norm in ["post", "pre"]:
    row = []
    for L in (16, 64, 256):
        _, _, dp = simulate(norm, L)
        row.append(f"L={L}: {dp[0] / dp[-1]:.3f}")
    print(f"  {norm}-Norm  " + ", ".join(row))

运行输出:

post-Norm, L=64:  前向范数/√d 1.00, 1.00, 1.00, 1.00, 1.00 | 隐状态梯度 0.56, 0.59, 0.54, 0.68, 0.82 | 参数梯度 8.9, 9.3, 8.8, 10.7, 12.5
  输入端/输出端参数梯度比 = 0.715

pre-Norm, L=64:  前向范数/√d 1.24, 3.22, 4.07, 4.93, 5.66 | 隐状态梯度 4.73, 1.94, 1.40, 1.13, 1.00 | 参数梯度 77.1, 30.8, 23.0, 17.3, 16.6
  输入端/输出端参数梯度比 = 4.651

深度敏感性(输入端/输出端参数梯度比):
  post-Norm  L=16: 0.853, L=64: 0.715, L=256: 0.167
  pre-Norm  L=16: 2.209, L=64: 4.651, L=256: 11.164

三笔账一目了然。前向:Post-Norm 的主干范数被 LN 钉死在 1.00;Pre-Norm 的主干范数随深度增长(L=64 时到 5.66,√64=8 的量级)——这正是 Xiong 等引理 2 的形态,也解释了 Pre-LN 顶层为什么必须有 final LN 来重标定。反向:Pre-Norm 的隐状态梯度靠恒等主干直传,从顶层 1.00 到底层 4.73,分支贡献沿途按多项式温和累积;Post-Norm 各层都在 0.5~0.8 之间徘徊,每层都付了收缩税。关键在深度敏感性:同样加深度,Post-Norm 的输入端/输出端参数梯度比按指数恶化(0.853 → 0.715 → 0.167),输入端收到的梯度随深度指数缩水;Pre-Norm 的比值虽然也在涨(2.21 → 4.65 → 11.16),但只是 √L 量级的多项式失衡,且恒等主干保证梯度必然到达每一层。

需要说明:这是线性近似的放大率模型,只考察初始化时刻,不替代对真实 Transformer 的完整分析;但「Post-Norm 的账随深度指数变坏、Pre-Norm 只是多项式失衡」的形态与论文的理论和实验一致。

稳与优的调和,以及 QK-Norm 的新动向

「Pre-Norm 更稳但略伤性能」是工程界多年的经验——Nguyen 与 Salazar(2019)在机器翻译里就观察到 Pre-Norm 训练更稳、BLEU 却略低。调和方案走两条路。一条把 Post-Norm 拉回牌桌:微软 DeepNet(2022)的 DeepNorm 把残差乘上大于 1 的系数 α、初始化按 β 缩小,让 Post-LN 稳定训到 1000 层且免 warmup。另一条接受 Pre-Norm,再往注意力内部加归一化:Gemma 3(2025)移除了与 FlashAttention 内核不兼容的 attention logit softcapping(tanh 截断),改用 QK-Norm——对 Q、K 向量做 RMSNorm;Qwen3 的 q_norm/k_norm 同样是在 RoPE 之前对 head 维做 RMSNorm,还顺带修掉了 Qwen2 在 fp16 推理时的注意力溢出问题;Qwen3-Next 又把 RMSNorm 权重的初始化从 1 改成 0 以进一步稳住训练。归一化从「层与层之间」一路内卷到了「注意力内部的头维度」。

边界在哪

归一化管尺度,不管内容。 它保证每层输入的统计尺度稳定,不保证信息不丢。RMSNorm 论文自己的 CIFAR-10 数据里,BatchNorm 错误率 8.25% 仍是最佳,RMSNorm 8.83%,LayerNorm 10.49%——在批内统计天然同分布的视觉 CNN 里,BN 至今仍是强基线。归一化方式的选择跟着数据结构走,不是越新越好。

省掉中心化是有条件的。 上面演示里,均值近零时两种归一化几乎重合(差 0.019),均值抬高 10 后差 3.7——当某层激活天然强正偏(例如 ReLU 之后全正)时,中心化与否的差别就不是小量。RMSNorm 的「效果不降」是在大模型的典型激活分布下成立的。

推理一致性是隐藏红利,数值精度是新软肋。 LN/RMSNorm 没有统计量要缓存,训练与推理同构,这是它胜过 BN 的工程原因;但注意力 logits 在 fp16 下溢出的现实案例(Qwen2 在端侧的教训)说明,归一化的覆盖范围之外仍有稳定性死角——这正是 QK-Norm 登场的背景。

归一化位置与残差是整体设计。 Post 换 Pre 不是局部替换:warmup、初始化、学习率要整套重调;DeepNorm 的 α 与 β 也是成对设计。读现代模型的结构图时,与其记「用了哪个归一化」,不如记住这笔账的算法规:统计量来自哪里、归一化挂在残差哪一侧、梯度因此走什么通道——这三个决定,写下了整个训练稳定性的账本。

参考资料

  1. Layer Normalization — Ba, Kiros, Hinton (arXiv 2016)
  2. Root Mean Square Layer Normalization — Zhang & Sennrich (NeurIPS 2019)
  3. On Layer Normalization in the Transformer Architecture — Xiong et al. (ICML 2020)
  4. DeepNet: Scaling Transformers to 1,000 Layers — Wang et al. (2022)
  5. Gemma 3 Technical Report — Google DeepMind (2025)
  6. Qwen3 Technical Report (2025)
  7. Transformers without Tears — Nguyen & Salazar (IWSLT 2019)
← 返回资讯列表

读者留言

COMMENTS 暂无
仅本站原创文章开放留言 · 请勿留下手机号、邮箱等个人信息

还没有留言,来说第一句?