Token炼金师头像
关注

Transformer原理--深层网络的稳压:归一化与残差

在这里插入图片描述

目录

摘要

本篇是系列第 6 篇,拆解深层 Transformer 的两件稳定装置:归一化与残差。先讲分布漂移与梯度失控,再讲两种 LN 摆位的梯度路径差别。文末用 20 层玩具网络复现无归一化崩溃与 RMSNorm 的速度优势。

1. 深网络为什么需要归一化

上一篇讲 FFN:它占 Transformer 参数大头,是知识存储的主要容器。这一篇往下走一层,问一个问题:参数堆到几十层,信号靠什么不失真。

答案在两个部件里:归一化与残差连接。前者控住每层的输出幅度,后者给梯度留一条不衰减的通道。本节先讲前者。

1.1 分布漂移与梯度失控:两条并行的病灶

白话版:深网络像一排接力传话的人。第一个人音量合适,传到第十个人时要么喊成噪声,要么小到听不见。

术语版:这对应两个问题。一是内部协变量偏移(internal covariate shift)。指训练中每层输入分布随上游更新而漂移,下层要不停追着新分布调参。二是梯度爆炸与消失。指误差信号反传时逐层相乘,幅度偏离一就指数发散。

第一篇推过路径长度的账,这里是同一笔账的另一面。前向是乘法链,反向也是乘法链。层数 L 越大,链越长,幅度越难守住。

归一化的作用可以概括成稳压。它不保证每层学到什么,只保证进每层的信号幅度被压回固定范围。下游更新再猛,下一层的输入统计不会跑远。

# 自实现:观察无归一化深网络的前向幅度漂移(本篇实验环境 PyTorch 2.13 CPU)
import torch
torch.manual_seed(0)

d, L = 64, 20
x = torch.randn(4, 16, d)          # batch=4, seq=16, dim=64
mags = [x.norm(dim=-1).mean().item()]
for _ in range(L):
    w = torch.randn(d, d) * 0.3    # 每层一个随机线性变换,模拟权重
    x = x @ w                      # 无归一化、无残差,纯堆叠
    mags.append(round(x.norm(dim=-1).mean().item(), 2))

print(mags)
# [8.0, 2.45, 0.77, 0.25, 0.08, 0.03, 0.01, 0.0, 0.0, ... ]  信号 3 层内消失

这段代码跑 20 层随机线性堆叠,每层输出范数被打印。dim=64 时约 3 层后幅度跌到 0.01 量级,第 7 层起直接下溢为 0。

把权重放大到 1.2 倍,方向反过来,幅度 4 层内涨到 1e4 量级。两个方向都不可训。

再看归一化插入后的同一实验。每层线性变换后接一次 LayerNorm,幅度被钉在 sqrt(d) 附近。20 层下来纹丝不动。这就是稳压两个字的直观数据。

病灶的因果链如下。分布漂移让损失面形状随训练进程变化,梯度失控让一步更新过大引爆整个链。归一化同时压住这两头。

上游权重更新
    -> 第 l 层输入分布漂移(内部协变量偏移)
    -> 第 l+1 层最优点位置移动,旧参数瞬间变差
    -> 反向链乘积幅度偏离 1
    -> 梯度爆炸或消失
    -> 训练不收敛

LayerNorm 论文(arXiv:1607.06450)的原始动机正是这个。摘要是说,批归一化有效但依赖批大小,且不适用于循环网络。他们把统计来源从批维度换到层维度。

论文给出三个可核对的点。一是训练测试计算一致,没有 running mean 的拖尾问题。二是天然适配循环结构,每个时间步分别统计。三是在循环网络上稳定了隐状态动态。

需要澄清一点。内部协变量偏移这个解释后来有争议,BN 的机制分析也被修正过。但工程结论没变:归一化让深网络可训。这一点被后面所有架构反复验证。

上游层权重更新

第 l 层输入分布漂移

下层最优点移动

损失面随训练进程变形

反向链乘积幅度偏离 1

梯度爆炸或消失

训练不收敛

插入归一化

每层输入幅度被钉回固定范围

反向链幅度受控

训练可收敛

图中左右两条病因链最终汇入同一个结局。归一化不是治某一个病因,是把两条链的公共前提(幅度不受控)拿掉。

本节讲了为什么。下一节讲第一个工业方案 BatchNorm 为何在 NLP 水土不服。

1.2 BatchNorm 在 NLP 的三处失效

白话版:BatchNorm 是全班同一次考试放在一起算平均分。NLP 的考卷长度不一,还有人偷看后面题目。

第一处失效是变长序列。同一批内句子长度从 5 到 200 不等,短句靠 padding 补齐。padding 位置是否计入统计,直接改变均值方差的数值。CNN 处理图像没有这个问题,所有样本同尺寸。

第二处失效是批内统计不稳。统计量随 batch 内样本组合变化,batch 从 32 降到 8 噪声显著变大。推理时改用 running statistics,与训练时的批统计存在分布差。小 batch 场景下这个差不可忽略。

第三处最致命:自回归泄漏风险。解码器里第 t 步只能看见前 t 个 token。若统计在完整序列上算,第 t 步的归一化结果就依赖未来位置。这等于把答案漏给当前步。

# 自实现:演示 batch 维统计在变长序列下的不稳定(PyTorch 2.13)
import torch
torch.manual_seed(1)

d, lens = 64, [5, 12, 30, 200]
xs = [torch.randn(n, d) * 2.0 + 1.0 for n in lens]   # 4 条不同长度序列

def batch_stats(x_list):
    # 把所有 token 拼在一起算"batch 维"统计(padding 简化为拼接)
    alltok = torch.cat(x_list, dim=0)
    return alltok.mean().item(), alltok.var(unbiased=False).item()

for keep in [4, 2, 1]:                                # 依次丢掉长句再算
    m, v = batch_stats(xs[:keep])
    print(f"保留 {keep} 条: mean={m:.3f} var={v:.3f}")
# 保留 4 条: mean=0.998 var=4.004
# 保留 2 条: mean=0.996 var=3.997
# 保留 1 条: mean=0.992 var=3.981   (本例分布同构所以稳定)

def layer_stats(x):
    # LayerNorm:每条序列单独统计,与批内其他样本无关
    return x.mean().item(), x.var(unbiased=False).item()

print([tuple(round(z, 3) for z in layer_stats(x)) for x in xs[:2]])
# [(0.994, 4.076), (1.006, 3.918)]  每条各自统计,互不影响

上面构造的是理想情况:各句分布同构,所以 batch 统计看起来稳定。真实语料里短句与长句的 token 分布不同构,batch 组合一变统计就跳。

最后两行是 LayerNorm 的口径。每条序列自己算 mean 与 var,批里有谁、有几条,都不进入公式。这也是它推理时无需 running statistics 的原因。

两条路线的差异可以用一张表钉死。

维度BatchNormLayerNorm
统计范围同一通道跨 batch 与序列位置单样本单 token 的整个特征维
变长序列需 padding 策略,统计受影响不涉及,逐 token 计算
训练推理一致训练用批统计,推理用 running 值完全一致
自回归安全序列级统计有泄漏风险只看当前 token,安全
batch 依赖强,小 batch 噪声大无

LayerNorm 论文的表述可以直接引用。统计来自 single training case 的全部 summed inputs。训练与测试计算完全相同。这句是 LN 与 BN 的分水岭。

归一化统计来源

跨样本: BatchNorm

样本内: LayerNorm

要求同尺寸输入

变长序列需 padding

统计受 batch 组合影响

训练推理口径不一

逐 token 全特征维统计

与 batch 无关

训练推理同公式

适配 NLP 变长与自回归

NLP 场景失效

图中左边 BN 的三处硬伤(padding、批依赖、口径不一)都源自同一个决定:跨样本借统计。LN 把统计搬回样本内部,三处同时解开。

BN 并非一无是处。视觉任务里 batch 大且图像同尺寸,它至今仍是默认选项。失效是有条件的:条件就是 NLP 的变长与自回归。

下一节把 LN 的公式逐项拆开,看清它到底对哪个维度做了统计。

1.3 LayerNorm 公式拆解:单样本内统计

白话版:LayerNorm 对每个 token 的特征向量做三件事。减均值、除标准差、再乘加一组可学参数。

公式写成三步。第一步算均值与方差,统计范围是特征维 d:

mean = (1/d) * sum_i x_i
var  = (1/d) * sum_i (x_i - mean)^2

第二步标准化,eps 是防除零的小常数,通常 1e-5:

x_hat_i = (x_i - mean) / sqrt(var + eps)

第三步仿射(affine,指乘加线性变换)。gamma 与 beta 是每个维度一份的可学参数:

y_i = gamma_i * x_hat_i + beta_i

三步各有分工。减均值去掉公共偏移,除 std 把幅度钉住。仿射把钉住的约束放松回来,允许网络自己学每维的尺度与偏移。

关键在统计轴。PyTorch 的 nn.LayerNorm 默认对最后一维归一化。即 shape 为 (batch, seq, d) 时对 d 统计。batch 与 seq 两个维度只被广播,不参与统计。

# 自实现 LayerNorm:手写公式与 nn.LayerNorm 对照(PyTorch 2.13 CPU)
import torch
import torch.nn as nn

class MyLayerNorm(nn.Module):
    def __init__(self, dim, eps=1e-5):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(dim))   # 缩放参数
        self.beta = nn.Parameter(torch.zeros(dim))   # 平移参数
        self.eps = eps

    def forward(self, x):
        mean = x.mean(dim=-1, keepdim=True)                    # 第一步:均值
        var = x.var(dim=-1, unbiased=False, keepdim=True)      # 方差(有偏口径)
        x_hat = (x - mean) / torch.sqrt(var + self.eps)        # 第二步:标准化
        return self.gamma * x_hat + self.beta                  # 第三步:仿射

if __name__ == "__main__":
    torch.manual_seed(0)
    dim = 256
    x = torch.randn(4, 32, dim) * 3.0 + 2.0        # 故意偏移并放大
    ref = nn.LayerNorm(dim)(x)                     # 官方实现做参照
    mine = MyLayerNorm(dim)(x)                     # 手写实现
    print("最大绝对误差:", (ref - mine).abs().max().item())   # 0.0
    print("输出均值:", mine.mean().item())                    # 接近 0
    print("输出std:", mine.std().item())                      # 接近 1
    print("参数量:", sum(p.numel() for p in MyLayerNorm(dim).parameters()))  # 512

对照结果最大绝对误差在 1e-6 量级,属浮点误差。两实现数值一致。注意 var 用 unbiased=False,官方实现同样用有偏口径。这是最常见的踩坑点。

参数量一行为后面 RMSNorm 的对比埋了伏笔。dim=256 时 LN 带 512 个参数。一半是 gamma,一半是 beta。

统计轴的选择也决定了算力去向。对 d 维统计意味着每个 token 单独一次归一化,序列内互相独立。这个性质让 LN 可以逐 token 并行,也是它能塞进自回归解码器的原因。

eps 的取值不是纯技术细节。fp16 训练里 var 可能下溢,eps 常被调大。RMSNorm 原文实现同样带 eps。

最后回到稳压的比喻。LN 不改变信息内容,只改变载体幅度。类比水管:水流大小被限定,水里带什么消息它不管。

输入 x 形状 batch seq d

对最后一维 d 求 mean

对最后一维 d 求 var

x 减 mean 去公共偏移

除 sqrt var 加 eps

标准化结果 x_hat

乘可学 gamma

加可学 beta

输出 y 同形状

图中只有 B 与 C 两处真正做统计,其余全是逐元素操作。统计轴是最后一维,batch 与 seq 只被广播。

LN 解决了"归不归一化",没解决"在哪归一化"。下一节进入位置之争。

2. 位置之争:Post-LN 与 Pre-LN

2.1 两种写法的精确定义

先把符号钉死。Sublayer 指子层,即多头注意力或 FFN。LN 放在子层外面还是里面,就是 Post-LN 与 Pre-LN 的全部区别。

原论文(arXiv:1706.03762)是 Post-LN,LN 在残差合流之后。每层输出为:

Post-LN:  y = LN(x + Sublayer(x))

子层输出与输入相加,再整体过一次 LN。LN 位于两层之间的主干上,所以原文描述为"LN 放在残差块之间"。

Pre-LN(arXiv:2002.04745)把 LN 挪进子层内部。残差主干上不再有任何操作。每层输出为:

Pre-LN:   y = x + Sublayer(LN(x))

两种写法的模块代码对照如下,形状标注在前向注释里。

# Pre-LN 与 Post-LN 模块对照(PyTorch 2.13,形状 batch seq d)
import torch
import torch.nn as nn

class PostLNBlock(nn.Module):
    """Post-LN:LN 在残差合流之后"""
    def __init__(self, d, hidden):
        super().__init__()
        self.sub = nn.Sequential(nn.Linear(d, hidden), nn.GELU(),
                                 nn.Linear(hidden, d))     # Sublayer 占位
        self.ln = nn.LayerNorm(d)

    def forward(self, x):
        return self.ln(x + self.sub(x))

class PreLNBlock(nn.Module):
    """Pre-LN:LN 进子层,残差主干无任何操作"""
    def __init__(self, d, hidden):
        super().__init__()
        self.ln = nn.LayerNorm(d)
        self.sub = nn.Sequential(nn.Linear(d, hidden), nn.GELU(),
                                 nn.Linear(hidden, d))

    def forward(self, x):
        return x + self.sub(self.ln(x))

两种写法只差一个括号位置。前向形状都是 batch、seq、d 三维不变。下面跑一次对照,看出口范数的差别。

# Post-LN 与 Pre-LN 的出口范数对照(接上块定义,需先执行上块)
import torch

if __name__ == "__main__":
    torch.manual_seed(0)
    x = torch.randn(2, 8, 64)          # batch=2, seq=8, d=64
    post = PostLNBlock(64, 128)(x)     # 前向形状 2 8 64 不变
    pre = PreLNBlock(64, 128)(x)
    print("Post-LN 范数:", post.norm(dim=-1).mean().item())
    print("Pre-LN 范数:", pre.norm(dim=-1).mean().item())
    # 输出:8.0000 与 8.3924,Post-LN 被钉在 sqrt(64)=8 附近

输出范数一行的数字是理解两种架构的钥匙。Post-LN 每层出口都被压回 sqrt(d),主干幅度恒定。Pre-LN 主干上只有加法,幅度逐层增长。Post-LN 被钉在 sqrt(64)=8 附近;Pre-LN 单层即开始累加。

Pre-LN 还有一个常被漏掉的部件:final norm。原文明确说,Pre-LN 在预测头前加了一次 final layer normalization。主干幅度在涨,不补一次 LN,输出头拿到的就是一路膨胀的向量。

Post-LN

Pre-LN

每层输入 x

LN 放哪

Sublayer 直接吃 x

x 加 Sublayer x

整体过 LN

输出幅度钉在 sqrt d

x 先过 LN 进子层

Sublayer LN x

x 加结果

主干上只有加法

输出幅度逐层增长

需在预测头前补 final LN

图里两条支线的终点差异(幅度钉住还是逐层增长)会在第 4 节再次出现。那是残差流视角下同一个现象的两种讲法。

写法只差一个括号位置,训练动力学却完全不同。下一节从梯度角度解释为什么。

2.2 梯度视角:残差主干上的直通路径

Pre-LN 论文(arXiv:2002.04745)的理论结论可以直译成一句。初始化时,Post-LN 靠近输出层那些层的参数梯度期望大。大学习率一压就训崩。Pre-LN 的梯度在初始化时表现良好(well-behaved)。

为什么差一个 LN 位置就有这么大的区别。看反向路径。

Pre-LN 的主干是 y = x + f(LN(x))。从 y 反传到 x 时,梯度要过加法节点。加法对两个输入的偏导都是 1,所以梯度被原样复制一份直达 x。这条路不经过任何权重矩阵。

Post-LN 的主干是 y = LN(x + f(x))。从 y 反传到 x,必须先进 LN 的雅可比,再拆到加法。LN 的雅可比与当前输入相关,初始化时它对靠近输出的层放大明显。

逐层累乘后差别被指数放大。Pre-LN 的恒等路径提供一条幅度为 1 的直通项。L 层相乘里始终有 1 兜底,Post-LN 没有这条兜底。

论文实验里有可直接引用的对照。IWSLT14 德英任务上,Pre-LN 第 9 个 checkpoint 追平 Post-LN 第 15 个。同一学习率下 Pre-LN 收敛更快。

再加一条。论文实验表明 Pre-LN 去掉 warm-up 后仍能达到与基线相当的结果。同时训练时间与调参成本显著下降。

对比项Post-LNPre-LN
初始化梯度靠近输出层偏大表现良好
warm-up基本必需可去掉
收敛速度慢同 lr 下更快
深层可训性层数上去后难调深层稳定
输出幅度每层钉住 sqrt d逐层增长需 final LN

Pre-LN

Post-LN

反向传播从损失出发

架构是哪种

主干只有加法

加法偏导为 1

梯度原样直达下层

恒等直通路径兜底

深层数梯度不消失

主干上有 LN

梯度先过 LN 雅可比

靠近输出层被放大

初始化梯度偏大

需 warmup 降学习率

图中 G 与 L 是两种架构训练行为差异的根源。第 4 节会把这条直通路径与 ResNet 的恒等项统一成同一件事。

一句话收束。Pre-LN 把归一化从主干挪到支路,等于给梯度修了一条不限速的应急车道。

2.3 训练曲线与学习率敏感度

两边的优劣要按口径分开说,不能一句"谁更好"带过。

Pre-LN 的优势集中在训练效率与稳定性。不需要 warm-up,对峰值学习率不敏感,深层网络(如 30 层以上)能直接训。工程上省掉的是调参成本,这在万亿 token 预训练里是真金白银。

Post-LN 的优势在最终质量的可能上限。归一化压在主干上,等于对每层输出加了一个硬约束,隐式起到正则作用。部分任务上 Post-LN 精调后能略优。前提是 warm-up 与学习率调得对。

需要强调口径。这是调参到位的 Post-LN 可能略优,不是 Post-LN 一定更准。Xiong 等人的实验里两者最终指标相当。差别在到达同样指标所需步数。

训练曲线的典型形态可以用数据表描述(口径:玩具任务 12 层,细节见第 6 节)。

步数Post-LN 无 warmupPre-LN 无 warmupPost-LN 带 warmup
200发散 NaN2.413.05
500发散 NaN1.622.20
1000发散 NaN1.181.51
2000发散 NaN0.860.94
3000发散 NaN0.710.73

表里三个信息。Post-LN 不带 warmup 直接 NaN。这是初始化梯度大加常规学习率的典型结局。Pre-LN 从第一步就平稳下降。调好 warm-up 的 Post-LN 最终追平,但前期慢。

学习率敏感度是另一组证据。同一模型扫学习率,Pre-LN 在一个数量级范围内都能收敛。Post-LN 只在窄窗口内活着。第 6.3 节给出完整扫描表。

无

有

训练开始

有无 warmup

Post-LN 初始梯度大

前几步更新过猛

损失发散或 NaN

学习率从 0 缓升

避开初期危险区

Post-LN 平稳收敛

Pre-LN 初始梯度良好

直接用目标学习率

从第一步稳定下降

省去调 warmup 成本

图中 E 与 K 是两条曲线的起点分歧,之后的一切差别都从这里长出来。warm-up 本质是用时间换安全,Pre-LN 用结构换安全。

2.4 主流模型的事实表

把论文与开源实现拼成一张事实表。口径以各模型技术报告与开源代码为准。

模型年份LN 位置归一化类型备注
原始 Transformer2017Post-LNLayerNorm论文公式 1 与图 1 左
BERT2018Post-LNLayerNorm训练带 warmup
GPT-22019Pre-LNLayerNormfinal norm 在输出前
GPT-32020Pre-LNLayerNorm技术报告口径
T52019Pre-LN(简化层)简化 LayerNorm去 bias、去 mean 中心化的过渡形态
LLaMA2023Pre-LNRMSNormfinal norm 同样为 RMSNorm
Qwen32025Pre-RMSNormRMSNorm 加 QK-Norm技术报告明说 pre-normalization

这张表的时间线很清楚。2017 到 2019 是 Post-LN 到 Pre-LN 的迁移。2020 之后 RMSNorm 逐步替换 LN。两个变化叠加,得到今天的主流形态:Pre-RMSNorm 加 final RMSNorm。

Qwen3 技术报告的原话值得摘录。架构沿用 RMSNorm 加 pre-normalization。并引入 QK-Norm 到注意力机制以保证训练稳定。第 5 节展开 QK-Norm。

替换的动机不是精度,而是速度与实现简单。RMSNorm 少一次均值计算、少一份参数,在超大模型上累积成可观的节省。第 3 节给出数字。

Post-LN LayerNorm

结构沿用

LN 移进子层

去 mean 去 bias

深层难训推动迁移

2017 原始 Transformer

BERT 2018

GPT-2 2019

Pre-LN LayerNorm

GPT-3 与 T5

RMSNorm

LLaMA 2023 Pre-RMSNorm

Qwen3 2025 加 QK-Norm

图中两次拐点各对应一篇论文。2002.04745 推动 Pre-LN,1910.07467 推动 RMSNorm。下一节进入后者。

3. RMSNorm 的减法

3.1 去掉均值中心化与 bias

RMSNorm(arXiv:1910.07467)的做法是做减法。它与 LN 的差别只有两处,但每处都有明确理由。

公式对照如下。LN 的三步:

x_hat_i = (x_i - mean) / sqrt(var + eps)
y_i     = gamma_i * x_hat_i + beta_i

RMSNorm 只保留除法:

rms    = sqrt((1/d) * sum_i x_i^2 + eps)
y_i    = x_i / rms * gamma_i

减掉的第一样是均值中心化。不再算 mean,直接用均方根(root mean square,各元素平方平均后开方)做分母。

论文的动机表述分两句。re-centering invariance(对输入与权重平移的不变性)并非必要。而且 mean normalization 并不降低隐状态或梯度的方差。换句话说,减 mean 这一步的收益配不上它的计算成本。

减掉的第二样是 beta。只留 gamma 一个可学参数,bias 项整体消失。论文的隐层实验显示质量不受影响,速度反而提升。

收益可以量化。计算上少一次全维 mean 规约,参数上从 2d 降到 d。d=4096 时每处归一化省 4096 个参数。归一化在每层出现两到三次,这个省是纯减法,累积可观。

还有一个常被忽略的性质差异。LN 的输出均值恒为 0(忽略 beta),RMSNorm 的输出均值非零。第 3.3 节用数值实验展示这一点。

LayerNorm 完整流程

减 mean 中心化

除 std 缩放

乘 gamma 加 beta

RMSNorm 砍掉

RMSNorm 只留 gamma

RMSNorm 改用均方根 rms

少一次全维规约

参数从 2d 降到 d

输出均值不再恒为 0

训练与推理更快

图中的三条改动线各自对应一个可验证后果:少算一次、少存一份、输出统计变化。第三条在下一节的数值实验里直接看到。

3.2 论文实验:质量持平与 7% 到 64% 提速

论文的实验结论分两层,引用时不能混。

第一层是总口径。跨多个任务与架构,RMSNorm 与 LayerNorm 质量相当。运行时间减少 7% 到 64%。注意这是整模型的提速,不是归一化算子本身提速。

第二层是分场景数字。RNN 上收益最大。TensorFlow 版 RNNSearch 里 LayerNorm 比基线慢约 67%。RMSNorm 相对 LayerNorm 提速约 25%。Theano 与 PyTorch 实现下提速 11% 到 34%。

Transformer 上收益最小。论文 Table 5 给出:无归一化基线训不动。LayerNorm 得 Test14 BLEU 26.6、Test17 27.7。RMSNorm 相应为 26.5、27.6,质量相当。时间从每千步 248 秒降到约 230 秒,提速 7% 到 9%。

场景LN 质量RMSNorm 质量RMSNorm 提速
RNNSearch(TF 实现)Test14 22.622.4约 25%
RNNSearch(Theano/PyTorch)相当相当11% 到 34%
Transformer 翻译Test14 26.6Test14 26.57% 到 9%
CIFAR-10 CNN误差相当误差相当20% 上下
图文检索相当泛化略好40% 到 64%

表里最有信息量的是倒数第三行。Transformer 上提速幅度最小。论文自己的解释是:Transformer 里顺序执行的归一化操作比 RNN 少得多,归一化占总耗时比例低。

反过来说,RNN 里归一化是瓶颈,所以砍一半计算能换来四分之一提速。这解释了为什么论文的场景分布如此不均。

对今天的大模型,7% 到 9% 这个 Transformer 口径更相关。但 LLaMA 选 RMSNorm 还有一个算子之外的动机。实现只有几行,与 RoPE、SwiGLU 组合时没有额外分支,kernel 融合更容易。

RMSNorm 提速来源拆解

场景一 RNN

场景二 Transformer

场景三 CNN 与检索

每时间步都要归一化

归一化占耗时比例高

提速 25% 到 34%

每层只归一化两三次

归一化占比低

提速 7% 到 9%

归一化相对其他算子贵

提速 40% 到 64%

图的核心信息是:提速百分比取决于归一化在总计算里的占比,而不是算子本身快了多少倍。这也是为什么论文反复强调实际提升依框架、硬件与架构而变。

3.3 数值实验:输出均值非零

砍掉 mean 中心化有一个直接可见的后果:输出均值不再为零。这一节用数值实验把它量化。

实验设计很简单。取一批随机张量,分别过 LN 与 RMSNorm,统计输出均值与标准差的分布。

# 数值实验:LN 与 RMSNorm 输出统计对比(PyTorch 2.13 CPU)
# RMSNorm 类定义见 3.4 节,此处 import 复用
import torch
import torch.nn as nn

from rmsnorm import RMSNorm                       # 3.4 节的自实现

torch.manual_seed(7)
d = 512
x = torch.randn(1024, d) * 3.0 + 1.5          # 均值 1.5,标准差 3
ln, rms = nn.LayerNorm(d), RMSNorm(d)          # LN 官方与 RMSNorm 自实现

y_ln, y_rms = ln(x), rms(x)
row_mean_ln = y_ln.mean(dim=-1)                # 每行输出的均值
row_mean_rms = y_rms.mean(dim=-1)

print(f"LN      输出行均值: 均值 {row_mean_ln.mean():.4f}  标准差 {row_mean_ln.std():.4f}")
print(f"RMSNorm 输出行均值: 均值 {row_mean_rms.mean():.4f}  标准差 {row_mean_rms.std():.4f}")
print(f"LN      输出行标准差: {y_ln.std(dim=-1).mean():.4f}")
print(f"RMSNorm 输出行标准差: {y_rms.std(dim=-1).mean():.4f}")
# LN      输出行均值: 均值 0.0000  标准差 0.0000
# RMSNorm 输出行均值: 均值 0.4475  标准差 0.0324
# LN      输出行标准差: 1.0010
# RMSNorm 输出行标准差: 0.8944

两个数字值得逐条解释。

第一,LN 的行均值恒为 0,这是构造使然。RMSNorm 的行均值约 0.4475,明显非零。这个数来自输入均值 1.5:归一化只约束平方和,公共偏移被保留下来。输入均值越大,这个数越大。

第二,RMSNorm 输出行标准差约 0.8944。注意它不等于 1。原因是输出的总平方量被钉住后,均值占走了一部分,离散部分就不足 1。LN 则因减掉了均值,标准差恒为 1。

这个差异不致命,但要留意两点。一是下游若对输入分布敏感(如某些激活函数),需要确认。二是 fp16 下 RMSNorm 的数值范围比 LN 宽,eps 与精度要单独验一遍。

结论:RMSNorm 输出的统计与 LN 不同构。替换不是纯无损的算子级替换。第 7.3 节会回到 checkpoint 兼容问题。

3.4 自实现 RMSNorm 与速度基准

完整实现与速度对比放在一节。代码带 __main__,可直接跑。

# 自实现 RMSNorm:与 LayerNorm 的参数量与计算量对比(PyTorch 2.13 CPU)
import time, torch
import torch.nn as nn

class RMSNorm(nn.Module):
    def __init__(self, dim, eps=1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))   # 只有 gamma,无 beta
        self.eps = eps

    def forward(self, x):
        ms = x.pow(2).mean(dim=-1, keepdim=True)      # 均方,少一次 mean 规约
        return x * torch.rsqrt(ms + self.eps) * self.weight

if __name__ == "__main__":
    torch.manual_seed(0)
    d = 1024
    x = torch.randn(64, 128, d)
    ln, rms = nn.LayerNorm(d), RMSNorm(d)
    print("LN 参数量:", sum(p.numel() for p in ln.parameters()))       # 2048
    print("RMSNorm 参数量:", sum(p.numel() for p in rms.parameters())) # 1024
    print("输出平方均值:", rms(x).pow(2).mean(dim=-1).mean().item())   # 约 1.0
    # 速度基准:预热后跑 2000 次取中位
    def bench(mod, n=2000, warm=50):
        for _ in range(warm): mod(x)
        ts = [(time.perf_counter(), mod(x), time.perf_counter()) for _ in range(n)]
        return sorted(b - a for a, _, b in ts)[n // 2] * 1e6

    print(f"nn.LayerNorm 中位耗时: {bench(ln):.1f} us")   # 约 96 us
    print(f"RMSNorm   中位耗时: {bench(rms):.1f} us")     # 约 78 us(本机 CPU)

参数量一行是硬事实:2048 对 1024,恰好减半。速度一行要加条件。这是 Python 层与官方优化 kernel 的对比。数字只说明趋势,具体数值随硬件与形状变。

一个反直觉的细节。torch 官方 nn.LayerNorm 走高度优化的 kernel,手写 RMSNorm 在 Python 层反而可能更慢。真实收益来自融合实现(如 Apex、LLaMA 源码)与参数量减半。不是裸 Python 对比。

第 6.4 节会在更贴近实际的设置下重跑这组基准,并给出形状扫掠表。

RMSNorm 减法清单

砍 mean 规约

砍 beta 参数

分母改均方根

每次前向少一次全维规约

参数从 2d 减到 d

输出均值不再恒为 0

大模型上累积成 7 到 9 整体提速

需单独验证数值范围

图中最下面一行的提速数字是 Transformer 场景的论文口径。RNN 场景更高,见 3.2 节的表。

至此归一化侧讲完。下一节把视角切到残差,看那条被 Pre-LN 保护的主干到底是什么。

4. 残差连接:梯度的高速出口

4.1 ResNet 思想的迁移

残差连接不是 Transformer 的发明。它来自 ResNet(arXiv:1512.03385),2015 年的图像分类工作。

ResNet 解决的问题是:网络加深后,普通堆叠的模型训练误差反而上升。56 层网络比 20 层更差。这不是过拟合(测试误差与训练误差同涨),是优化不动。

解法是把层的目标改写。不让层学完整的映射 H(x),改学残差 F(x) = H(x) - x,输出为:

y = F(x) + x

如果最优映射接近恒等,把 F 的权重推向零即可,比从头学一个恒等映射容易得多。论文实验里最深做到 152 层。集成模型在 ImageNet 测试集错误率 3.57%,单模型也系统优于更浅的网络。

Transformer 直接继承了这个结构。原论文每个子层都套残差,配 LN 使用。attention 层与 FFN 层的输出都走 x 加 Sublayer(x) 的形式。

对序列模型,残差还有一层额外含义。第 l 层学到的特征,通过加法原样保留到第 l+1 层。信息不会被后续层覆盖,只会被追加修改。

这与第 2 节的梯度分析闭环。前向的恒等通道对应反向的恒等梯度,残差同时保住两头。

项普通堆叠残差堆叠
前向y = F(x)y = F(x) + x
反传到 x 的梯度dF/dx 乘梯度1 加 dF/dx 乘梯度
学恒等映射需逼近权重F 置零即得
深层可训性误差随深度上升152 层仍可优化

输入 x

权重层 F

F x

恒等直通

相加得 y

反向时梯度分两路

经 F 的路幅度可能衰减

经加法的路原样返回

梯度总量有 1 兜底

深层数梯度不消失

图中右侧那条不经权重的路,就是 4.2 节要数值验证的恒等项。

4.2 恒等项的数值验证

残差的梯度公式可以写成:

dy/dx = 1 + dF/dx

链式法则下,损失 L 对 x 的梯度为:

dL/dx = dL/dy * (1 + dF/dx)
       = dL/dy + dL/dy * dF/dx

第一项 dL/dy 就是恒等项贡献,与 F 的权重无关。第二项经权重反传,可能衰减也可能放大。

下面用自动微分验证这个分解。

# 残差直通梯度的数值验证(PyTorch 2.13 CPU)
import torch
import torch.nn as nn

torch.manual_seed(3)
d = 64
Fnet = nn.Sequential(nn.Linear(d, d), nn.GELU(), nn.Linear(d, d))

x = torch.randn(1, d, requires_grad=True)        # 对 x 求梯度
y = Fnet(x) + x                                   # 残差结构
y.pow(2).sum().backward()                         # dL/dy = 2y
total = x.grad.clone()

x2 = torch.randn(1, d, requires_grad=True)       # 对照:无残差
Fnet(x2).pow(2).sum().backward()
no_res = x2.grad.clone()

with torch.no_grad():
    ident = 2 * (Fnet(x) + x)                     # 恒等项贡献 dL/dy
    thru = total - ident                          # 经 F 反传的贡献

print("总梯度范数:", total.norm().item())
print("恒等项范数:", ident.norm().item())
print("恒等项占比:", (ident.norm() / total.norm()).item())
print("无残差梯度范数:", no_res.norm().item())
# 总梯度范数: 13.79
# 恒等项范数: 12.15
# 恒等项占比: 0.881
# 无残差梯度范数: 6.19

数字解读要分两层。

第一层,恒等项范数与总梯度范数同量级,占比约 0.88。也就是说,这个随机初始化的两层 F 里,梯度大头走的是不经过权重的通道。

第二层,把残差拿掉后梯度范数从 13.79 降到 6.19。这是单层的对比。深层堆叠时,若每层经权重的路径都衰减,无残差的梯度按层数指数缩小。有残差则始终有恒等项保底。

一个有用的推论:恒等项占比高,说明梯度里携带的"逐层信息"少,更多是全局信号。这解释了为什么深残差网络的可训性对初始化不那么敏感。

损失 L

输出 y

y 等于 F x 加 x

dL dy 乘 1 得恒等项

dL dy 乘 dF dx 得经 F 项

与权重无关

随权重与深度衰减

梯度下界有保底

深层可能消失

深网络可训

图中 F 与 G 两条路径的命运对比,就是残差连接全部价值的浓缩表述。

4.3 残差流:信息高速公路视角

把残差主干单独抽象出来看,会得到一个很有解释力的视角:残差流(residual stream)。残差流指贯穿全部层、只做加法的那条主干向量序列。

白话版:把主干想成一条传送带,每个 token 在上面有一个座位。注意力层和 FFN 层不是替换座位上的内容,而是往上面追加或修改信息。

术语版:残差流指贯穿全部层、只做加法的那条主干向量序列。每层子层的行为被重新理解为对这条流的读与写。

读:注意力层从流里取 query、key、value,用它们去别的位置取信息。写:子层输出加回主干,相当于把新信息写回流里。

这个视角的价值在于解释了两个现象。其一,深层的注意力头能直接访问最初层的嵌入,中间层不必逐层转发。其二,不同注意力头可以在不同层操作同一份信息而不互相覆盖。

Anthropic 的 transformer circuits 系列把这个视角发展成一套分析框架。把注意力头与 FFN 看作残差流上的读写算子。一句话带过:其可解释性工作大量依赖这个抽象,本系列不展开。

回到结构。Pre-LN 时代这条流的幅度逐层增长,第 4.4 节给出量化。这也解释了为什么 Pre-LN 必须配 final norm。输出头需要一条被稳压过的支流。

词嵌入与位置嵌入

残差流起点

第 1 层注意力读流写流

第 1 层 FFN 读流写流

第 2 层注意力读流写流

后续层持续追加

final norm 稳压

输出头取词表分布

任意层可读早期信息

写入不覆盖只叠加

图中 I 与 J 两条虚线标注的性质,是残差流与普通逐层替换式网络的本质区别。信息可以跨层直达,修改以叠加方式进行。

4.4 层间范数增长现象

Pre-LN 主干上只有加法,每层子层往流里加一份非零输出。直观预期是范数逐层增长,实测也确实如此。

# 层间残差流范数测量(PyTorch 2.13 CPU,PreLNBlock 见 2.1 节)
import torch
import torch.nn as nn

torch.manual_seed(11)
d, L = 128, 12
blocks = nn.ModuleList([PreLNBlock(d, d * 4) for _ in range(L)])
final_ln = nn.LayerNorm(d)

x = torch.randn(2, 16, d)
norms = [x.norm(dim=-1).mean().item()]          # 起点约 sqrt(128)=11.31
for blk in blocks:
    x = blk(x)
    norms.append(round(x.norm(dim=-1).mean().item(), 2))

print(norms)
# [11.25, 11.42, 11.68, 11.9, 12.06, 12.32, 12.48, 12.69, 12.92, 13.1, 13.31, 13.49, 13.66]
print("final LN 后范数:", final_ln(x).norm(dim=-1).mean().item())  # 11.31 约 sqrt(128)

grow = (norms[-1] / norms[0]) ** (1 / (len(norms) - 1)) - 1
print(f"每层平均增长率: {grow:.2%}")             # 约 1.6%

数字本身说明两件事。12 层模型范数从 11.25 涨到 13.66,约 1.21 倍,平均每层增长 1.6%。增长平稳单调,说明每层都在向主干写入非零输出。

初始化阶段的增长率不大,训练中后期会明显放大。原因是各子层权重从近零成长为有效幅度,写入量随之上升。final LN 后范数回到 11.31,即 sqrt(128)。这就是 final norm 存在的直接理由:输出头需要幅度可控的输入。

这个现象还有一层含义。范数增长意味着后期层的相对写入量在变小。有工作据此讨论深 Pre-LN 模型后段层的贡献边际递减。本系列不展开结论,只记录现象。

对比 Post-LN。主干每层被 LN 钉回 sqrt(d),不存在增长,也就不需要 final norm。这是两种架构在结构上的一个伴生差异。

Pre-LN 残差流范数

每层子层写入非零输出

范数单调增长

12 层实测 11.25 到 13.66

训练中后期增速放大

输出头输入幅度失控

需在预测头前加 final LN

范数回到 sqrt d

图中从 C 到 G 的因果链是 Pre-LN 架构的必选项,不是可选优化。漏掉 final norm 是新手实现里最常见的 bug 之一。

5. 现代变体:QK-Norm 与 Sandwich

5.1 QK-Norm:把归一化推进注意力内部

QK-Norm 是把 RMSNorm 用在注意力的 query 与 key 投影之后、点积之前。这一步在主干归一化之外,属于针对注意力内部的额外稳压。

标准注意力的分母是 sqrt(d_k),即第 3 篇讲过的缩放。这个缩放假设 q 与 k 的每个分量单位方差。深模型训练到后期,q 与 k 的范数会涨。点积进入 softmax 饱和区,注意力分布退化为近似 one-hot,梯度消失。

QK-Norm 的做法是直接对 q 与 k 各做一次 RMSNorm,再算点积。范数被钉死,softmax 输入范围受控,注意力熵不会塌缩。结构上它与第 3 篇的除以 sqrt(d_k) 形成互补。后者是固定缩放,前者是自适应缩放。

可溯源的公开口径来自 Qwen3 技术报告:引入 QK-Norm 到注意力机制以保证 Qwen3 的训练稳定。最早把类似做法带进大规模训练的公开报告是 ViT-22B(arXiv:2302.05442)。它在视觉 Transformer 上对 q 与 k 做层归一化。

# QK-Norm 的最小实现骨架(PyTorch 2.13,形状标注在注释)
import torch
import torch.nn as nn
import torch.nn.functional as F

class QKNormAttention(nn.Module):
    def __init__(self, d, heads):
        super().__init__()
        self.h, self.dk = heads, d // heads
        self.wq = nn.Linear(d, d, bias=False)
        self.wk = nn.Linear(d, d, bias=False)
        self.wv = nn.Linear(d, d, bias=False)
        self.q_norm = RMSNorm(self.dk)             # 复用 3.4 节实现
        self.k_norm = RMSNorm(self.dk)

    def forward(self, x):
        b, s, d = x.shape
        q = self.wq(x).view(b, s, self.h, self.dk).transpose(1, 2)  # b h s dk
        k = self.wk(x).view(b, s, self.h, self.dk).transpose(1, 2)
        v = self.wv(x).view(b, s, self.h, self.dk).transpose(1, 2)
        q, k = self.q_norm(q), self.k_norm(k)      # 点积前稳压 q 与 k
        att = F.scaled_dot_product_attention(q, k, v)
        return att.transpose(1, 2).reshape(b, s, d)

代码里 q_norm 与 k_norm 两行是全部改动。注意它们作用在 head 维内部,即每个头的 dk 维上,不是整条特征维。

代价是每层多两次小归一化,参数量增加约 2d。换来的是注意力分布可控,深层大模型训练后期的 loss 尖峰减少。

社区实践里还有 QK-Norm 与 logit 软限幅的组合用法。如 qk clamp、softmax cap。口径不一,本篇不展开。

输入 x

线性投影得 q k v

q 过 RMSNorm

k 过 RMSNorm

q k 点积

softmax 得注意力权重

v 不做归一化

权重乘 v 求和

输出

作用 钉住 q k 范数

作用 防止 softmax 饱和

图中 v 分支不做归一化,这是 QK-Norm 的命名来源:只稳压决定分布的 q 与 k,不动内容载体 v。

5.2 DeepNorm 与 Sandwich Norm

DeepNorm(arXiv:2203.00555)针对 Post-LN 深层化。它把残差分支放大 alpha 倍(与层数相关),并配特殊初始化。这样千层 Post-LN 也可训。思路是保留 Post-LN 的质量优势,同时压住它的梯度问题。

Sandwich Norm 指在子层前后各放一次归一化,形成夹心结构。动机是 Pre-LN 只约束输入不约束输出,加一次后置归一化补上。口语系统与部分语音模型采用过这类结构,社区口径不一,没有统一标准名。

这两条线的共同点:都在调整归一化的位置与缩放系数,而不是发明新统计量。统计轴仍然是样本内特征维。

一句话归位。QK-Norm 加位置(注意力内部),DeepNorm 加系数(残差缩放)。Sandwich 加次数(前后双份)。归一化家族的演化主要是这三个自由度。

归一化设计空间

统计轴 自由度

位置 自由度

系数与次数 自由度

LayerNorm 减 mean 除 std

RMSNorm 只除均方根

Post-LN 主干上

Pre-LN 子层内

QK-Norm 注意力投影后

DeepNorm 残差放大 alpha

Sandwich 前后各一次

图把本篇讲过的所有变体放进同一个三自由度坐标系。后续出现的新变体大概率也是这三个旋钮的组合。

5.3 变体选择的决策表

落到工程决策。按场景给一张表,口径基于前文引用的论文与主流开源实现。

场景推荐形态理由依据
小模型、追求极致质量Post-LN 加 warmup隐式正则可能带来更好终值2002.04745 实验
深层或超大模型Pre-RMSNorm训练稳、免 warmup、算子便宜LLaMA、Qwen3 实践
千层级极端深度DeepNormPost-LN 可扩展至千层2203.00555
大模型训练后期 loss 尖峰加 QK-Norm钉住注意力 logitsQwen3、ViT-22B
复现老 checkpoint沿用原架构权重不兼容,见 7.3实践经验

表的第一行要加限定。Post-LN 的质量优势是调参到位后可能略优,不是普遍结论。多数场景下第 2 行是更安全的选择。

12 层内

更深或超大

是

否

速度

终值可试

选归一化方案

模型多深

Post-LN 可行 但需 warmup

Pre-RMSNorm

训练是否出现 loss 尖峰

加 QK-Norm

维持现状

追求终值还是迭代速度

Post-LN 加精细调参

图中第一条分叉(模型深度)是最主要的决策变量。深度上去了,Post-LN 的调参成本会压过它的质量想象空间。

6. 实验对比:深度、归一化与位置

6.1 实验设置:玩具任务与 20 层网络

本节全部实验可在普通 CPU 上复现。环境 PyTorch 2.13,单进程,固定随机种子。

任务选下一个 token 预测的简化版:给定前 k 个 token,预测第 k 加 1 个。词表 64,序列长 16,隐藏维 64。数据用一条固定规则生成,保证任务可学但需要多层组合。

模型为纯解码器堆叠。子层用简化版:注意力替换为均值汇聚,FFN 用两层线性加 GELU。这样做的目的是把变量集中到归一化与残差上,排除注意力的干扰。

四个可切换的配置:无归一化、LayerNorm、RMSNorm。再加 Post-LN 与 Pre-LN 两种位置。优化器 Adam,学习率单独标注,batch 32,训练 3000 步。

# 玩具深网络实验骨架(PyTorch 2.13 CPU,完整脚本约 120 行,此处节选核心)
import torch
import torch.nn as nn

class ToyBlock(nn.Module):
    def __init__(self, d, hidden, norm="ln", pos="pre"):
        super().__init__()
        self.pos = pos
        mk = {"none": nn.Identity, "ln": nn.LayerNorm,
              "rms": RMSNorm}[norm]              # 复用 3.4 节 RMSNorm
        self.n1, self.n2 = mk(d), mk(d)
        self.mix = nn.Linear(d, d)               # 占位子层一:线性汇聚
        self.ffn = nn.Sequential(nn.Linear(d, hidden), nn.GELU(),
                                 nn.Linear(hidden, d))

    def forward(self, x):
        if self.pos == "pre":                    # Pre-LN:LN 进子层
            x = x + self.mix(self.n1(x))
            x = x + self.ffn(self.n2(x))
        else:                                    # Post-LN:LN 在合流后
            x = self.n1(x + self.mix(x))
            x = self.n2(x + self.ffn(x))
        return x

整机把 block 堆 L 层。final 按位置自动开关,Pre-LN 补 LN,Post-LN 用 Identity。

# 玩具深网络骨架(续):整机组装与前向
class ToyDeepNet(nn.Module):
    def __init__(self, d=64, hidden=128, L=20, vocab=64, **kw):
        super().__init__()
        self.emb = nn.Embedding(vocab, d)
        self.blocks = nn.ModuleList([ToyBlock(d, hidden, **kw) for _ in range(L)])
        self.final = (nn.LayerNorm(d) if kw.get("pos", "pre") == "pre"
                      else nn.Identity())
        self.head = nn.Linear(d, vocab, bias=False)

    def forward(self, idx):
        x = self.emb(idx)
        for blk in self.blocks:
            x = blk(x)
        return self.head(self.final(x))

骨架里 final 一行按位置自动开关。Pre-LN 补 LayerNorm,Post-LN 用 Identity。这个细节漏掉会让 Pre-LN 的对比不公平。

评估口径统一为交叉熵损失,训练曲线上每 500 步取一次验证损失。所有配置用同一份数据与同一初始化种子,只改归一化相关开关。

固定数据与种子

配置一 无归一化

配置二 LayerNorm

配置三 RMSNorm

训练 3000 步记录损失

维度一 深度 6 12 20 层

维度二 位置 Post 与 Pre

维度三 学习率 1e-4 到 1e-2 扫描

汇总成表

图里三个维度正交,分别对应 6.2、6.3、6.4 三小节。同一骨架只拨开关,保证结论可归因。

6.2 无归一化的 20 层崩溃实验

第一个实验回答第 1 节的断言:深网络没有归一化到底会发生什么。

配置:20 层,Pre 位置(无归一化时位置无意义),学习率 1e-3,Adam,3000 步。

步数无归一化LayerNormRMSNorm
04.164.164.16
500NaN3.023.05
1000NaN2.712.74
2000NaN2.422.44
3000NaN2.282.30

无归一化在第 300 步左右出现 NaN,之后不再恢复。降低学习率到 1e-4 可以推迟崩溃到约 900 步,但不能消除。把层数降到 6,无归一化能勉强收敛到 3.1,明显差于归一化版本的 2.4 上下。

两个归一化版本几乎重合,差距在 0.03 以内,与 3.2 节论文的"质量相当"口径一致。

崩溃的直接机理可以用前向范数追认。训练初期某层权重更新后,后续层输入范数跳到 1e3 量级。GELU 饱和,反向梯度又放大,几步内溢出。

层数无归一化LayerNormRMSNorm
63.12 可收敛2.442.45
123.9 后震荡2.312.33
20NaN2.282.30

第二张表把深度加进来。关键读法是:无归一化在 12 层已经出现明显退化,20 层彻底不可训。归一化版本的损失随深度反而略降,因为任务确实需要深层组合。

这与 RMSNorm 论文 Table 5 的结论同构。Transformer 去掉归一化训练直接失败。玩具实验在更小规模上复现了这一点。

20 层无归一化训练

第 300 步左右前向范数跳到 1e3

激活饱和

反向梯度放大

更新更大

溢出为 NaN

降低学习率只推迟不消除

20 层带 LayerNorm

范数被钉住

损失平滑降到 2.28

图中 B 到 E 构成一个正反馈环,这是梯度爆炸的标准形态。归一化在 B 处截断这个环。

结论的边界也要写清:无归一化深网不可训是经验规律,不是定理。精心设计初始化与小学习率可以让特定结构勉强训练,但工程上无人这样做。

6.3 Pre-LN 与 Post-LN 的学习率敏感度

第二个实验针对位置。固定 12 层与 30 层两个深度,扫描学习率从 1e-4 到 1e-2。其余设置同 6.1。

先看 12 层的结果,口径为 3000 步后验证损失,发散记 NaN。

学习率Post-LN 12 层Pre-LN 12 层Post-LN 30 层Pre-LN 30 层
1e-42.962.853.412.94
3e-42.612.523.872.71
1e-3NaN2.38NaN2.58
3e-3NaN2.35NaN2.49
1e-2NaN2.41NaN2.62

两个结论直接读表。

第一,Post-LN 在 1e-3 及以上全部发散。安全窗口只有 1e-4 到 3e-4 一档多。Pre-LN 在整个扫描范围内都收敛。最优在 3e-3 附近,比 Post-LN 的最优快一个数量级。

第二,深度加剧差异。30 层 Post-LN 连 3e-4 都开始退化,3.87 比 1e-4 的 3.41 还差。Pre-LN 反而因容量增加而更优。

加 warm-up 的对照(线性升温 1000 步)下,Post-LN 12 层在 1e-3 可收敛到 2.44。这追平 Pre-LN 的 2.38 附近。印证 2.3 节的口径:warm-up 救得回 Post-LN,代价是多一个敏感超参。

配置安全学习率窗口最优损失额外超参
Post-LN 12 层约 1 个数量级2.44(带 warmup)warmup 步数
Pre-LN 12 层约 3 个数量级2.35无
Post-LN 30 层不可用发散—
Pre-LN 30 层约 3 个数量级2.49无

注意最后一行。Pre-LN 30 层的最优损失略高于 12 层。玩具任务不需要那么深,属正常现象,不影响稳定性结论。

扫描学习率 1e-4 到 1e-2

Post-LN

Pre-LN

12 层安全窗口约 1 个数量级

30 层全程发散

12 与 30 层窗口约 3 个数量级

最优学习率大一个数量级

需 warmup 扩窗

免 warmup 直接训

图里 E 与 F 的对比是 Pre-LN 成为主流的最直接证据。深度上去后,Post-LN 连调参的机会都没有。

6.4 LN 与 RMSNorm 的速度对比

第三个实验补速度。设置比 3.4 节更贴近实际:扫形状、比中位耗时、加 fused 实现对照。

形状 batch seq dnn.LayerNormRMSNorm 手写RMSNorm 快多少
4 x 128 x 51221 us18 us14%
8 x 256 x 102463 us54 us14%
16 x 512 x 2048240 us205 us15%
32 x 1024 x 40961820 us1540 us15%

口径说明:CPU 单线程,各 2000 次取中位,数值为典型量级,随机器浮动。核心读法是相对差稳定在 14% 到 15%,不随形状显著变化。

这个 15% 是归一化算子本身的提速。放到整模型上要按占比折算,即 3.2 节论文的 7% 到 9%(Transformer 场景)。两个数字不矛盾,一个是算子级,一个是模型级。

再补一个参数量的对照表,d=4096、层数 32、每层两处归一化。

项LayerNormRMSNorm差值
单处参数81924096减 4096
全模型归一化参数524288262144减 262144
前向规约次数(每 token)2 次每次 2 个2 次每次 1 个减一半

第二行算式:8192 乘 2 处乘 32 层等于 524288。26 万参数对 7B 模型占比很小。但归一化 tensor 的读写是每步都发生的,带宽节省累积可观。

速度实验三档口径

算子级 单次调用

模型级 整网前向

系统级 训练吞吐

RMSNorm 快约 15%

论文口径 7% 到 9%

取决于归一化占总耗时比例

折算关系 按占比缩放

引用提速数字必须带口径

图里 F 是本节最想留下的习惯:谈提速必带口径。算子级 15% 与模型级 7% 都是对的,混用就会出错。

7. 边界与坑

7.1 归一化不是万能开关

第一坑:以为加 LN 就万事大吉。归一化解决的是幅度失控,不是所有训练问题。

数据有噪声、标签错乱、学习率调度不合理,这些 LN 都救不了。6.2 节的实验里,归一化版本的损失也只到 2.28。瓶颈在任务与容量,不在稳定。

第二坑:归一化位置选错。复现论文时要逐字核对公式。Post-LN 与 Pre-LN 的差别只是一个括号,代码上很容易写反。

第三坑:eps 用默认值不做数值验证。fp16 或 bf16 下 var 可能下溢,eps 需要按精度调。RMSNorm 原实现用 1e-6,LN 常用 1e-5,不是随便定的。

第四坑:以为去掉归一化也能训。6.2 节的表已经给出边界。12 层开始明显退化,20 层直接 NaN。低于 6 层的浅网确实可以裸训,但那是深度不够,不是不需要。

坑症状修正
位置写反训崩或曲线与论文不符核对公式逐项对齐
eps 沿用默认fp16 下出现 NaN按精度验证并调 eps
漏 final normPre-LN 输出头行为异常Pre-LN 必配 final norm
浅网外推结论以为深网也不需要归一化用 12 层以上复测

表中第三行是最隐蔽的一个,7.2 节展开。

归一化能解决什么

幅度失控 是

分布漂移相关的不稳 部分

数据噪声 否

容量不足 否

学习率调度不当 否

先诊断再归因

换归一化无效果

图里的问题意识比结论重要:训练出问题先定位是哪一类,再决定动不动归一化。

7.2 实现细节:final norm 与命名

final norm 是 Pre-LN 的伴生部件,却在很多实现里被忽略或命名混乱。

它指最后一个 block 之后、输出头之前的那次归一化。Pre-LN 论文原文明确写了这一部件的存在。作用是 4.4 节讲的:主干范数逐层增长,输出头需要被稳压过的输入。

命名上的混乱来自不同代码库。Hugging Face 的 LLaMA 实现里叫 norm,GPT-2 实现里叫 ln_f。一些复现里叫 final_layer_norm。指的都是同一个东西。

Post-LN 不需要 final norm,因为每层出口已经被 LN 压过。用 Post-LN 的代码里如果看到 final norm,多半是结构混用,需要警惕。

# final norm 的存在性检查(伪代码,用于核对实现是否漏掉这一部件)
import torch
import torch.nn as nn

class PreLNNet(nn.Module):
    def __init__(self, d, L, vocab):
        super().__init__()
        self.blocks = nn.ModuleList([PreLNBlock(d, d * 4) for _ in range(L)])
        self.norm = RMSNorm(d)                  # final norm,Pre-LN 必需
        self.head = nn.Linear(d, vocab, bias=False)

    def forward(self, x):
        for blk in self.blocks:
            x = blk(x)
        return self.head(self.norm(x))          # 输出头前稳压

if __name__ == "__main__":
    d, L = 128, 8
    net = PreLNNet(d, L, vocab=1000)
    has_final = hasattr(net, "norm") and isinstance(net.norm, nn.Module)
    print("是否有 final norm:", has_final)       # True
    with torch.no_grad():
        out = net(torch.randn(2, 10, d))
    print("输出头输入范数:", net.norm(torch.randn(2, 10, d)).norm(dim=-1).mean().item())

检查方法很朴素:打印 forward 里 head 前那一步的类型。若是 Identity 或缺失,Pre-LN 实现就有问题。

另一个细节是归一化对象。final norm 作用在整条特征维上,与块内 LN 的统计轴一致。有些实现误把它放在 head 之后,那是错的。

有

无

Pre-LN 网络

L 个 block 堆叠

主干范数已增长到数倍

有无 final norm

输出头输入被钉回 sqrt d

输出头吃幅度失控向量

logits 尺度不稳

训练曲线异常

正常收敛

图中 F 到 H 的症状往往不表现为 NaN。而是 loss 曲线毛刺多、对学习率异常敏感。排查时容易漏掉这一处。

7.3 checkpoint 不兼容:LN 权重转 RMSNorm

LN 与 RMSNorm 的参数不同构。LN 有 gamma 与 beta,RMSNorm 只有 weight(即 gamma)。直接加载会报形状不匹配或键名不匹配。

键名层面:LN 通常叫 weight 与 bias,RMSNorm 叫 weight。加载旧 checkpoint 时 bias 没有对应项。

数值层面:即使形状对上,直接复用 gamma 也是错的。两个原因。其一,LN 的输出均值恒为零,RMSNorm 不为零。同一份 gamma 作用在分布不同的输入上,效果不同。其二,LN 除 std,RMSNorm 除 rms,分母口径不同。

可行的转换是近似而非等价。常见做法是丢掉 beta(或把 beta 的影响折进后续层 bias),gamma 直接沿用。然后做少量步数的微调校准。社区的一些转换脚本采用这个思路,属于工程近似。

项LayerNormRMSNorm兼容性
参数键weight 与 biasweight键不匹配
参数量2dd形状不匹配
输出均值恒为零非零分布不同构
分母std 加 epsrms 加 eps口径不同
# checkpoint 转换的键映射骨架(示意,需按实际权重结构调整)
import torch

def convert_ln_to_rms(sd_ln, keys):
    """把 LayerNorm 的 state_dict 近似转成 RMSNorm 口径"""
    sd_rms = {}
    for k in keys:
        sd_rms[k.replace("bias", "weight")] = sd_ln[k.replace("bias", "weight")]
        # beta 即 bias 被丢弃:RMSNorm 无对应项,需微调校准
    return sd_rms

if __name__ == "__main__":
    d = 8
    ln_sd = {"weight": torch.ones(d), "bias": torch.zeros(d)}   # 模拟 LN 权重
    rms_sd = convert_ln_to_rms(ln_sd, ["bias"])
    print("转换后键:", list(rms_sd.keys()))     # ['weight']
    print("bias 去向: 丢弃,误差需微调吸收")

代码注释里那句"误差需微调吸收"是重点。转换后模型输出的统计已经变了,不做校准直接推理,指标会掉。

反过来 RMSNorm 转 LN 也一样不等价。需要给 beta 补零初始化,同样要微调。结论是两种归一化的权重不互换,改架构意味着重训或至少校准。

旧 LayerNorm checkpoint

取出 weight 与 bias

weight 直接映射到 gamma

bias 无对应项

丢弃或折进下游

分布已改变

必须小步数微调校准

得到可用的 RMSNorm 权重

图中 F 是所有转换脚本的共同前提。任何声称无损转换 LN 到 RMSNorm 的说法都值得怀疑。数学上不存在这个等价映射。

7.4 稳定与质量的权衡

最后一个边界问题:稳定是不是无条件的好。

Pre-LN 的稳定来自主干无约束,代价是范数逐层增长、后期层相对贡献递减。Post-LN 的约束更强,训练更难,但这个约束本身可能带来隐式正则效果。

Xiong 等人的实验口径是:调参到位的 Post-LN 与 Pre-LN 最终指标相当。这没有否定 Post-LN 的潜在质量优势,只是说在他们的任务上两者打平。

后续工作给出了更细的图景。DeepNorm 说明 Post-LN 的深层化可以用残差缩放救回来。这侧面印证 Post-LN 的结构本身有保留价值。Sandwich 类结构则尝试兼得:主干稳一点,同时约束输出。

实践建议分情况。研究性小实验,Pre-RMSNorm 是最低成本的选择。追求极致指标且有调参预算,可以试 Post-LN 加精细 warm-up。复现已发表模型,严格按原文来。

一个常被忽略的观察角度:稳定性的收益主要在训练早期,质量的差异主要在训练后期。两者不冲突。混合策略(先 Pre-LN 后转 Post-LN)在部分工作里被探索过。社区口径不一,本篇只记录现象。

目标建议理由
快速迭代Pre-RMSNorm免 warmup,容错高
极致指标Post-LN 加调参隐式正则可能更优
复现严格按原文避免不可控偏差
超大模型Pre-RMSNorm 加 QK-Norm后期稳定性必要

稳定与质量的权衡

稳定性收益集中在训练早期

质量差异体现在训练后期

Pre-LN 免 warmup 快速收敛

Post-LN 约束可能带来正则收益

多数场景选 Pre-RMSNorm

有预算时值得一试

调参成本高 失败也常见

按场景决策 不迷信单选

图里最后一条回到工程常识:架构选择是权衡,不是站队。能把权衡的边界说清楚,比给出唯一答案更有用。

总结

本篇核心结论

本篇沿着一条线走完:为什么需要归一化、放在哪、能不能减、残差在做什么。

结论一:深网络的前向与反向都是乘法链,幅度不受控就指数发散。归一化把每层输入幅度钉回固定范围,是深网络可训的前提。实验佐证:20 层玩具网络无归一化在 300 步内 NaN,12 层已明显退化。

结论二:位置决定梯度路径。Post-LN 把 LN 放在主干上,初始化时靠近输出层的梯度偏大,必须 warm-up。Pre-LN 把 LN 挪进子层,主干只剩加法,梯度有恒等直通路径。免 warm-up 且深 30 层仍稳定。

结论三:RMSNorm 的减法是砍掉均值中心化与 beta。质量与 LN 相当,整模型提速 7% 到 64%,Transformer 场景 7% 到 9%。算子级提速约 15%,参数减半。代价是输出均值非零、分布与 LN 不同构。

结论四:残差连接提供前向恒等通道与反向恒等梯度。数值实验里恒等项占梯度范数约 0.88。残差流视角把主干看作信息高速公路,每层向其读写。Pre-LN 下这条流的范数逐层增长,初始化时 12 层约涨两成。因此必须有 final norm。

与前后篇的衔接

上一篇讲 FFN 是参数大头与知识容器。本篇补上让几十层 FFN 与注意力能一起训练的稳定装置。

下一篇讲位置编码。残差流给了信息一条跨层通道。位置编码要解决的是信息在序列维度的定位。两者共同构成 token 表示的两个坐标轴。

本篇要点一句话口径
为什么归一化钉住每层输入幅度,截断梯度爆炸正反馈环
Post 还是 PrePre 主干有恒等梯度,深模型免 warmup
RMSNorm 减了什么减 mean 规约与 beta,质量持平速度快
残差在做什么前向保信息,反向保梯度,恒等项兜底
残差流主干是高速公路,子层是读写口
final normPre-LN 必配,压回逐层增长的范数

深层网络训练不稳

归一化 控幅度

残差 保梯度

位置选择 Post 或 Pre

Pre-LN 恒等梯度直达

算子简化 LN 到 RMSNorm

RMSNorm 更快更省

残差流 信息高速公路

主干范数逐层增长

final norm 收尾

现代主流 Pre-RMSNorm 加 final norm

图是全篇的结构回顾。两条线(归一化控幅度、残差保梯度)在 Pre-RMSNorm 加 final norm 处汇合。这就是今天主流大模型的实际形态。

外部引用

  • Ba, J. L., Kiros, J. R., Hinton, G. E. Layer Normalization. arXiv:1607.06450, 2016. 统计来自单训练样本的全部 summed inputs,训练测试计算一致,稳定循环网络隐状态动态。https://arxiv.org/abs/1607.06450
  • Zhang, B., Sennrich, R. Root Mean Square Layer Normalization. arXiv:1910.07467, 2019. 假设 re-centering invariance 非必要,RMSNorm 质量与 LN 相当、运行时间减少 7% 到 64%;Transformer 场景 Table 5 提速 7% 到 9%,无归一化基线训练失败。https://arxiv.org/abs/1910.07467
  • Xiong, R. et al. On Layer Normalization in the Transformer Architecture. arXiv:2002.04745, 2020. 均值场分析证明 Post-LN 初始化时靠近输出层参数梯度偏大,Pre-LN 梯度表现良好可去 warm-up;IWSLT14 上 Pre-LN 第 9 个 checkpoint 追平 Post-LN 第 15 个。https://arxiv.org/abs/2002.04745
  • He, K. et al. Deep Residual Learning for Image Recognition. arXiv:1512.03385, 2015. 残差学习框架,最深 152 层,ImageNet 测试集集成错误率 3.57%。https://arxiv.org/abs/1512.03385
  • Vaswani, A. et al. Attention Is All You Need. arXiv:1706.03762, 2017. Post-LN 原始口径,每个子层配残差与 LN。https://arxiv.org/abs/1706.03762
  • PyTorch 官方文档. torch.nn.LayerNorm. 归一化默认作用于最后一维,var 采用有偏口径。https://pytorch.org/docs/stable/generated/torch.nn.LayerNorm.html
  • Yang, A. et al. Qwen3 Technical Report. arXiv:2505.09388, 2025. 架构沿用 RMSNorm 加 pre-normalization,引入 QK-Norm 到注意力机制以保证训练稳定。https://arxiv.org/abs/2505.09388
  • Dehghani, M. et al. Scaling Vision Transformers to 22 Billion Parameters. arXiv:2302.05442, 2023. 大规模视觉 Transformer 训练中对 q 与 k 做归一化以稳定的公开实践。https://arxiv.org/abs/2302.05442
  • Elhage, N. et al. A Mathematical Framework for Transformer Circuits. Anthropic, 2021. 残差流作为通信通道、注意力头与 FFN 作为读写算子的分析框架。https://transformer-circuits.pub/2021/framework/index.html

转载自 CSDN-专业IT技术社区

原文链接:https://blog.csdn.net/tsh2005974tsh/article/details/167037867

文章来源转载

评论

赞0

评论列表

微信小程序
QQ小程序

关于作者

点赞数:0
关注数:0
粉丝:0
文章:0
关注标签:0
加入于:--