AI 工程基础体系 · 第 62/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

神经网络初始化与归一化:Xavier、Kaiming、BatchNorm 和 LayerNorm

神经网络训练开始时,参数通常是随机初始化的。初始化过小,信号和梯度会在多层传播后逐渐消失;初始化过大,激活值或梯度可能爆炸。归一化则在训练过程中重新调整中间表示的尺度,使优化问题更容易处理。

初始化和归一化解决的是相关但不同的问题:

  • 初始化决定训练开始时参数和激活值的统计尺度。
  • 归一化在前向传播中根据当前数据重新调整激活值。
  • Xavier 初始化主要在输入输出尺度之间折中,常用于线性层、tanh 等近似对称激活函数。
  • Kaiming 初始化针对 ReLU 类激活函数补偿其截断负半轴造成的方差损失。
  • BatchNorm按批次统计特征,训练和推理阶段使用不同的统计来源。
  • LayerNorm按单个样本的特征维度统计,训练和推理通常使用相同的计算方式。

理解这些方法,需要先从信号如何穿过网络开始。


一、为什么初始化会影响训练稳定性

设某一层是线性变换:

zj=i=1nwjixi+bjz_j=\sum_{i=1}^{n}w_{ji}x_i+b_j

其中:

  • xix_i 是输入;
  • wjiw_{ji} 是权重;
  • bjb_j 是偏置;
  • nn 是输入维度,也称为 fan_in
  • zjz_j 是激活函数之前的预激活值。

先作几个常见近似:

  1. 输入和权重均值为 0;
  2. 权重和输入相互独立;
  3. 不同输入维度之间近似独立;
  4. 偏置初始为 0。

于是:

Var(zj)=i=1nVar(wjixi)\operatorname{Var}(z_j) = \sum_{i=1}^{n}\operatorname{Var}(w_{ji}x_i)

在独立条件下:

Var(wjixi)=Var(wji)Var(xi)\operatorname{Var}(w_{ji}x_i) = \operatorname{Var}(w_{ji})\operatorname{Var}(x_i)

因此:

Var(zj)=nVar(w)Var(x)\operatorname{Var}(z_j) = n\operatorname{Var}(w)\operatorname{Var}(x)

如果希望每层输出的方差大致保持不变,即:

Var(z)Var(x)\operatorname{Var}(z)\approx \operatorname{Var}(x)

就需要:

Var(w)1n\operatorname{Var}(w)\approx \frac{1}{n}

这只是线性层的前向传播条件。实际网络还包含非线性激活,而且反向传播还要求梯度尺度不要快速变大或变小。

1. 深层网络中的方差连乘

假设每一层都使信号方差乘以常数 cc,经过 LL 层后:

Var(xL)cLVar(x0)\operatorname{Var}(x_L)\approx c^L\operatorname{Var}(x_0)

c<1c<1 时,方差指数衰减;当 c>1c>1 时,方差指数增长。即使 c=0.9c=0.9,经过 50 层后也只有:

0.9500.00520.9^{50}\approx 0.0052

这就是“每层只缩小一点”仍然会导致深层信号消失的原因。

梯度也有类似问题。若每层的雅可比矩阵都使梯度范数平均乘以 cgc_g,则:

x0cgLxL\|\nabla x_0\| \approx c_g^L\|\nabla x_L\|

初始化方法的目标不是让每个样本、每个神经元的值完全相同,而是让整体统计尺度在训练初期处于合理范围。


二、激活函数会改变方差传播

初始化不能脱离激活函数讨论。设:

xl+1=ϕ(zl)x_{l+1}=\phi(z_l)

其中 ϕ\phi 是激活函数。

1. ReLU 会丢弃一半负值

ReLU 定义为:

ReLU(z)=max(0,z)\operatorname{ReLU}(z)=\max(0,z)

如果 zz 服从均值为 0 的对称分布,大约一半值为负,会被置为 0。对于零均值高斯变量 zz,有:

E[ReLU(z)2]=12E[z2]\mathbb{E}[\operatorname{ReLU}(z)^2] = \frac{1}{2}\mathbb{E}[z^2]

因此,从“二阶矩”角度看,ReLU 使尺度大约减少一半。

需要注意,ReLU 输出的均值通常大于 0,所以它的方差并不严格等于输入方差的一半。若 zN(0,σ2)z\sim\mathcal{N}(0,\sigma^2),则:

E[ReLU(z)]=σ2π\mathbb{E}[\operatorname{ReLU}(z)] = \frac{\sigma}{\sqrt{2\pi}}

Var(ReLU(z))=σ22σ22π\operatorname{Var}(\operatorname{ReLU}(z)) = \frac{\sigma^2}{2} - \frac{\sigma^2}{2\pi}

工程推导中常用“二阶矩约减半”的近似,因为它足以导出 Kaiming 初始化的主要尺度。

2. Sigmoid 和 Tanh 还可能进入饱和区

Sigmoid 为:

σ(z)=11+ez\sigma(z)=\frac{1}{1+e^{-z}}

其导数为:

σ(z)=σ(z)(1σ(z))\sigma'(z)=\sigma(z)(1-\sigma(z))

zz 很大或很小时,Sigmoid 接近 1 或 0,导数接近 0,反向梯度会消失。

Tanh 的导数为:

ddztanh(z)=1tanh2(z)\frac{d}{dz}\tanh(z)=1-\tanh^2(z)

z|z| 较大时,Tanh 也会饱和。Xavier 初始化的一个重要目标,就是让线性输出的尺度不要过大,以便更多值处于激活函数的有效梯度区域。


三、Xavier 初始化:在前向和反向之间折中

Xavier 初始化也称为 Glorot 初始化。其核心思想是同时考虑:

  • 前向传播中输入维度 fan_in
  • 反向传播中输出维度 fan_out

设某层有:

  • 输入维度 ninn_{\text{in}}
  • 输出维度 noutn_{\text{out}}

前向传播若要保持方差,大致需要:

Var(w)1nin\operatorname{Var}(w)\approx\frac{1}{n_{\text{in}}}

反向传播若要保持梯度方差,大致需要:

Var(w)1nout\operatorname{Var}(w)\approx\frac{1}{n_{\text{out}}}

两者折中得到:

Var(w)=2nin+nout\boxed{ \operatorname{Var}(w)=\frac{2}{n_{\text{in}}+n_{\text{out}}} }

这就是 Xavier 正态初始化的常见形式:

wN(0,2nin+nout)w\sim\mathcal{N} \left( 0, \frac{2}{n_{\text{in}}+n_{\text{out}}} \right)

若使用均匀分布:

wU(a,a)w\sim U(-a,a)

均匀分布的方差为 a2/3a^2/3,令其等于上面的目标方差:

a23=2nin+nout\frac{a^2}{3} = \frac{2}{n_{\text{in}}+n_{\text{out}}}

得到:

a=6nin+nout\boxed{ a=\sqrt{\frac{6}{n_{\text{in}}+n_{\text{out}}}} }

因此 Xavier 均匀初始化为:

wU(6nin+nout,6nin+nout)w\sim U\left( -\sqrt{\frac{6}{n_{\text{in}}+n_{\text{out}}}}, \sqrt{\frac{6}{n_{\text{in}}+n_{\text{out}}}} \right)

1. 一个完整计算例子

考虑一个线性层:

nn.Linear(4, 2)

因此:

nin=4,nout=2n_{\text{in}}=4,\qquad n_{\text{out}}=2

Xavier 初始化的目标方差为:

Var(w)=24+2=13\operatorname{Var}(w) = \frac{2}{4+2} = \frac{1}{3}

均匀分布边界为:

a=66=1a = \sqrt{\frac{6}{6}} = 1

所以:

  • Xavier 正态初始化:标准差为 1/30.577\sqrt{1/3}\approx0.577
  • Xavier 均匀初始化:从 [1,1][-1,1] 中采样。

如果输入每个维度的方差约为 1,那么线性层输出的方差近似为:

Var(z)=4×13×1=43\operatorname{Var}(z) = 4\times\frac{1}{3}\times1 = \frac{4}{3}

这并不等于 1,因为 Xavier 同时照顾了前向和反向条件,而不是只满足前向条件。

2. Xavier 的适用范围

Xavier 通常适合:

  • 线性层;
  • Tanh;
  • 某些近似对称、不会大量截断负值的激活函数;
  • 没有明确使用 ReLU 类激活函数的浅层 MLP。

PyTorch 中可以这样使用:

import torch
from torch import nn

layer = nn.Linear(4, 2)

nn.init.xavier_uniform_(layer.weight)
nn.init.zeros_(layer.bias)

也可以指定激活函数对应的增益:

nn.init.xavier_uniform_(
    layer.weight,
    gain=nn.init.calculate_gain("tanh")
)

gain 会改变目标尺度。使用 Xavier 时,初始化方法、激活函数和 gain 应作为一个整体决定,而不是只机械地调用函数名中带有 xavier 的 API。


四、Kaiming 初始化:为 ReLU 补偿半波截断

Kaiming 初始化也称 He 初始化,专门针对 ReLU 及其近似变体设计。

对于 ReLU,前面已经得到:

E[ReLU(z)2]12E[z2]\mathbb{E}[\operatorname{ReLU}(z)^2] \approx \frac{1}{2}\mathbb{E}[z^2]

线性层有:

E[z2]ninVar(w)E[x2]\mathbb{E}[z^2] \approx n_{\text{in}}\operatorname{Var}(w)\mathbb{E}[x^2]

为了让 ReLU 后的二阶矩大致保持不变,需要:

12ninVar(w)1\frac{1}{2} n_{\text{in}} \operatorname{Var}(w) \approx1

因此:

Var(w)=2nin\boxed{ \operatorname{Var}(w) = \frac{2}{n_{\text{in}}} }

Kaiming 正态初始化为:

wN(0,2nin)w\sim\mathcal{N} \left(0,\frac{2}{n_{\text{in}}}\right)

对应标准差:

std(w)=2nin\operatorname{std}(w)=\sqrt{\frac{2}{n_{\text{in}}}}

Kaiming 均匀初始化的边界为:

a=6nina=\sqrt{\frac{6}{n_{\text{in}}}}

因为均匀分布 U(a,a)U(-a,a) 的方差是 a2/3a^2/3

1. 继续使用前面的层作为例子

对于 nn.Linear(4, 2)

nin=4n_{\text{in}}=4

Kaiming 正态初始化目标方差为:

Var(w)=24=0.5\operatorname{Var}(w) = \frac{2}{4} = 0.5

标准差为:

0.50.707\sqrt{0.5}\approx0.707

如果输入方差为 1,则线性输出的二阶矩近似为:

E[z2]4×0.5×1=2\mathbb{E}[z^2] \approx 4\times0.5\times1 = 2

经过 ReLU 后:

E[ReLU(z)2]12×2=1\mathbb{E}[\operatorname{ReLU}(z)^2] \approx \frac{1}{2}\times2 = 1

这就是 Kaiming 初始化的补偿过程:线性层先产生约两倍的二阶矩,ReLU 再将其约减半。

2. PyTorch 用法

import torch
from torch import nn

layer = nn.Linear(4, 2)

nn.init.kaiming_normal_(
    layer.weight,
    mode="fan_in",
    nonlinearity="relu",
)
nn.init.zeros_(layer.bias)

参数含义:

  • mode="fan_in":优先保持前向传播的激活尺度;
  • mode="fan_out":优先保持反向传播的梯度尺度;
  • nonlinearity="relu":告诉初始化函数使用 ReLU 对应的增益。

对于普通前馈网络,fan_in 是常见选择。某些以梯度传播为主要考虑的结构可能选择 fan_out,但这不是无条件更好。

3. Leaky ReLU 的差异

Leaky ReLU 定义为:

ϕ(x)={x,x0ax,x<0\phi(x)= \begin{cases} x,&x\ge0\\ ax,&x<0 \end{cases}

其中 aa 是负半轴斜率。由于负半轴不再完全置零,所需的初始化增益不同。PyTorch 可以通过 a 指定斜率:

layer = nn.Linear(128, 128)

nn.init.kaiming_normal_(
    layer.weight,
    a=0.01,
    mode="fan_in",
    nonlinearity="leaky_relu",
)

如果实际使用的是 nn.LeakyReLU(0.2),初始化时也应使用相同的负斜率。初始化假定的激活函数和实际激活函数不一致,会使理论尺度失效。


五、fan_infan_out 如何计算

对于二维线性层权重:

nn.Linear(in_features, out_features)

权重形状通常为:

[out_features, in_features]

因此:

fan_in=in_features\text{fan\_in}=in\_features

fan_out=out_features\text{fan\_out}=out\_features

卷积层还要乘以卷积核空间尺寸。对于二维卷积权重形状:

[out_channels, in_channels, kernel_height, kernel_width]

有:

fan_in=in_channels×kernel_height×kernel_width\text{fan\_in} = in\_channels \times kernel\_height \times kernel\_width

fan_out=out_channels×kernel_height×kernel_width\text{fan\_out} = out\_channels \times kernel\_height \times kernel\_width

例如:

nn.Conv2d(
    in_channels=3,
    out_channels=64,
    kernel_size=3,
)

其:

fan_in=3×3×3=27\text{fan\_in}=3\times3\times3=27

fan_out=64×3×3=576\text{fan\_out}=64\times3\times3=576

PyTorch 的初始化函数会根据参数张量形状计算这些值,但它对张量布局有约定。若自定义权重不是常见的线性层或卷积层布局,应确认 fan_infan_out 是否被正确解释。错误的布局会导致初始化尺度偏离预期。

一个常见反例是把矩阵转置后直接初始化:

weight = torch.empty(128, 64)
# 后续实际计算可能是 x @ weight,而不是 F.linear(x, weight)

如果实际计算语义与 PyTorch 默认 Linear 的权重布局不同,却仍按默认布局理解 fan_in,初始化尺度可能被交换。自定义层应明确写出:

z=xW还是z=Wxz=xW \quad\text{还是}\quad z=W x

再决定哪个维度是输入维度。


六、Xavier 与 Kaiming 的选择边界

可以用以下因果关系理解二者:

激活函数或结构 常见初始化起点 原因
Linear Xavier 或较简单的方差保持初始化 不存在截断效应
Tanh Xavier,并考虑 gain 需要控制进入饱和区的概率
ReLU Kaiming,nonlinearity="relu" ReLU 约丢弃一半二阶矩
Leaky ReLU Kaiming,并指定负斜率 负半轴仍保留部分信号
GELU、SiLU 没有一个完全等价的简单公式 激活函数不是硬半波截断,通常结合架构和实验验证

最后一行很重要:不能把所有非线性函数都简单归入“ReLU,所以使用 Kaiming”。GELU 和 SiLU 的输入输出统计与 ReLU 不同。Transformer 中常见的线性层,通常还会配合 LayerNorm、残差连接和特定的缩放策略,因此不能只靠一个初始化公式推断整个模型的稳定性。


七、归一化到底在做什么

归一化层通常先计算某个维度集合上的均值和方差:

μ=1mi=1mxi\mu=\frac{1}{m}\sum_{i=1}^{m}x_i

σ2=1mi=1m(xiμ)2\sigma^2=\frac{1}{m}\sum_{i=1}^{m}(x_i-\mu)^2

然后标准化:

x^i=xiμσ2+ϵ\hat{x}_i= \frac{x_i-\mu}{\sqrt{\sigma^2+\epsilon}}

最后应用可学习的仿射变换:

yi=γix^i+βiy_i=\gamma_i\hat{x}_i+\beta_i

其中:

  • ϵ\epsilon 是防止除零的小常数;
  • γ\gamma 是可学习的缩放参数;
  • β\beta 是可学习的偏移参数。

归一化并不是简单地“把数据压缩到 0 到 1”。它通常使指定维度上的均值接近 0、方差接近 1,然后通过 γ,β\gamma,\beta 允许网络恢复所需的尺度和偏移。

初始化与归一化的差异是:

  • 初始化只在参数创建时发生一次;
  • 归一化在每次前向传播时作用于激活;
  • 归一化的统计维度决定它改变了哪些样本之间的关系;
  • 归一化参数本身也会进入优化和模型检查点。

八、BatchNorm:按批次统计特征

BatchNorm 的典型输入形状取决于层类型。

对于 BatchNorm1d,常见输入为:

[N, C]

或:

[N, C, L]

其中:

  • NN 是 batch size;
  • CC 是通道或特征数;
  • LL 是额外的序列长度。

BatchNorm 通常对每个通道 cc 独立统计,统计维度包括 batch 维以及可能的空间或序列维度。对于二维图像输入 [N, C, H, W]BatchNorm2d 对每个通道统计 N,H,WN,H,W 这些维度。

对于固定通道 cc,训练时可以抽象为:

μB,c=1mk=1mxk,c\mu_{B,c} = \frac{1}{m} \sum_{k=1}^{m}x_{k,c}

σB,c2=1mk=1m(xk,cμB,c)2\sigma_{B,c}^{2} = \frac{1}{m} \sum_{k=1}^{m}(x_{k,c}-\mu_{B,c})^2

yk,c=γcxk,cμB,cσB,c2+ϵ+βcy_{k,c} = \gamma_c \frac{x_{k,c}-\mu_{B,c}} {\sqrt{\sigma_{B,c}^{2}+\epsilon}} +\beta_c

这里的 mm 是当前 BatchNorm 实际统计到的元素数量。

1. BatchNorm 的训练状态和推理状态

BatchNorm 有两套统计来源:

训练模式

model.train()

通常使用当前 mini-batch 的均值和方差,并更新运行统计量:

  • running_mean
  • running_var

推理模式

model.eval()

通常不再使用当前输入 batch 的统计量,而使用训练阶段累积的运行统计量。

因此,train()eval() 不只是控制 Dropout。对于 BatchNorm,它们改变了前向计算的统计来源。

一个最小示例:

import torch
from torch import nn

torch.manual_seed(0)

bn = nn.BatchNorm1d(3)
x = torch.tensor([
    [1.0, 10.0, 100.0],
    [2.0, 20.0, 200.0],
    [3.0, 30.0, 300.0],
])

bn.train()
y_train = bn(x)

print("训练输出的列均值:", y_train.mean(dim=0))
print("训练输出的列方差:", y_train.var(dim=0, unbiased=False))
print("运行均值:", bn.running_mean)

bn.eval()
y_eval = bn(x)

print("推理输出的列均值:", y_eval.mean(dim=0))

在训练输出中,各列均值通常接近 0,方差通常接近 1,因为当前 batch 直接参与了归一化。切换到 eval() 后,输出不一定在当前 batch 上均值为 0,因为它使用的是历史运行统计量。

2. track_running_stats 的影响

PyTorch 的 BatchNorm 默认跟踪运行统计量。若关闭运行统计量跟踪:

bn = nn.BatchNorm1d(3, track_running_stats=False)

训练和推理时都依赖当前输入统计量。这会使推理结果依赖输入 batch 的组成,通常不适合需要单样本、固定输出或严格可复现的在线服务。

这不是 API 错误,而是统计语义发生了变化。部署前必须明确服务是:

  • 单样本推理;
  • 固定大小 batch 推理;
  • 动态 batch 推理;
  • 是否允许不同请求互相影响输出。

3. BatchNorm 的边界

BatchNorm 的统计依赖 batch,因此存在几个真实边界:

  1. batch 太小:均值和方差估计噪声大。
  2. 分布式训练:每张 GPU 只看到本地 batch,局部统计可能与全局统计不同。
  3. 数据分布变化:运行统计量来自历史训练数据,部署数据发生漂移时可能失配。
  4. 变长序列和 padding:若 padding 位置参与统计,统计量可能被无效位置污染。
  5. 训练和推理模式错误:忘记 eval() 会使推理输出受当前 batch 影响。
  6. 状态保存不完整:只保存参数而不保存 BatchNorm 的运行统计量,恢复后行为可能改变。

在多卡训练中,如果模型确实需要跨设备共享 batch 统计,通常要考虑同步 BatchNorm 等机制;但同步会引入跨设备通信成本,也可能降低吞吐。是否使用它取决于 batch 大小、模型结构和训练系统,而不是仅凭名称选择。


九、LayerNorm:按单个样本的特征维度统计

LayerNorm 不依赖 batch 中的其他样本。对输入最后若干个维度计算均值和方差。

对于形状:

[N, T, D]

如果使用:

nn.LayerNorm(D)

那么对每个样本、每个时间位置的 DD 个特征计算统计量:

μn,t=1Dd=1Dxn,t,d\mu_{n,t} = \frac{1}{D}\sum_{d=1}^{D}x_{n,t,d}

σn,t2=1Dd=1D(xn,t,dμn,t)2\sigma_{n,t}^{2} = \frac{1}{D}\sum_{d=1}^{D} (x_{n,t,d}-\mu_{n,t})^2

yn,t,d=γdxn,t,dμn,tσn,t2+ϵ+βdy_{n,t,d} = \gamma_d \frac{x_{n,t,d}-\mu_{n,t}} {\sqrt{\sigma_{n,t}^{2}+\epsilon}} + \beta_d

LayerNorm 的 normalized_shape 决定归一化哪些尾部维度。

import torch
from torch import nn

x = torch.randn(2, 5, 8)  # batch=2, sequence=5, hidden=8

norm = nn.LayerNorm(8)
y = norm(x)

print(y.shape)  # torch.Size([2, 5, 8])

print(y[0, 0].mean())                    # 通常接近 0
print(y[0, 0].var(unbiased=False))       # 通常接近 1

这里的统计只针对 y[0, 0, :] 这 8 个特征。它不会把两个样本、五个时间位置混在一起。

1. 错误的 normalized_shape

如果输入是 [N, T, D],但写成:

nn.LayerNorm(T)

通常会触发形状错误,因为 LayerNorm 默认匹配输入的最后一个维度,而最后一个维度是 D,不是 T

如果希望对最后两个维度 [T, D] 一起归一化,则应写:

nn.LayerNorm((T, D))

但这会把序列长度固定为 T,变长输入可能无法直接使用。Transformer 中使用 nn.LayerNorm(D),通常是因为每个 token 的隐藏维度 D 才是需要归一化的特征维度。

2. LayerNorm 为什么适合 Transformer

Transformer 的输入常见形状为:

[batch, sequence_length, hidden_size]

序列长度可能变化,batch size 也可能很小,甚至在线推理时为 1。LayerNorm 对单个 token 的 hidden features 统计,因此:

  • 不依赖 batch size;
  • 不需要运行均值和运行方差;
  • 训练和推理通常使用同一计算路径;
  • 适合自回归生成和动态 batch 服务。

LayerNorm 仍然有计算成本:它需要对每个 token 的 hidden 维度进行归约,并应用可学习的缩放和偏移。hidden size 很大时,这部分成本不能忽略,但它不会像 BatchNorm 那样引入跨样本统计依赖。


十、BatchNorm 与 LayerNorm 的统计维度差异

假设输入为:

x.shape = [N, T, D]

可以通过一个具体矩阵理解两者的区别。

BatchNorm 风格的统计

如果把每个特征 dd 作为一个通道,BatchNorm 可能在 N,TN,T 维度上统计:

μd=1NTn=1Nt=1Txn,t,d\mu_d = \frac{1}{NT} \sum_{n=1}^{N}\sum_{t=1}^{T}x_{n,t,d}

它回答的是:

在整个 batch 和序列位置中,第 dd 个特征的总体均值是多少?

LayerNorm 的统计

LayerNorm 对每个 (n,t)(n,t) 单独统计:

μn,t=1Dd=1Dxn,t,d\mu_{n,t} = \frac{1}{D} \sum_{d=1}^{D}x_{n,t,d}

它回答的是:

对这个样本的这个 token,所有 hidden features 的均值是多少?

因此两者不是同一种归一化的不同 API 名称,而是统计对象不同:

特性 BatchNorm LayerNorm
统计对象 同一特征跨样本,可能还跨空间或时间 单个样本内部的特征
是否依赖 batch
是否有运行统计量 通常有 没有 BatchNorm 式运行统计量
batch size 为 1 的稳定性 可能较差 通常不受影响
常见场景 CNN、视觉模型 Transformer、序列模型
推理行为 依赖 train/eval 状态 通常训练推理计算方式一致

十一、归一化层中的可学习参数与初始化

归一化层通常包含:

  • weight,对应 γ\gamma,初始为 1;
  • bias,对应 β\beta,初始为 0。

这样初始化后,归一化层一开始主要执行标准化,不额外改变尺度和均值:

yx^y\approx\hat{x}

例如:

from torch import nn

bn = nn.BatchNorm1d(128)
ln = nn.LayerNorm(128)

print(bn.weight.shape)  # [128]
print(bn.bias.shape)    # [128]
print(ln.weight.shape)  # [128]
print(ln.bias.shape)    # [128]

对于 BatchNorm,running_meanrunning_var 不是普通的梯度参数,通常不通过反向传播更新,而是作为模块状态在训练时更新。使用 state_dict() 保存模型时,应同时保存这些状态。

state = model.state_dict()
torch.save(state, "checkpoint.pt")

如果只手动保存 named_parameters() 得到的可训练参数,却遗漏 BatchNorm 的运行统计量,恢复后的推理结果可能与保存前不一致。


十二、初始化与归一化如何共同工作

一个常见误解是:

既然有 BatchNorm 或 LayerNorm,初始化就不重要了。

归一化会减弱前几层尺度错误的影响,但不能消除所有初始化问题。

1. 归一化之前的数值仍可能溢出

如果权重初始化过大,线性层可能先产生极大的 z

z=Wxz=W x

之后即使进入归一化层,也可能出现:

  • 激活函数在归一化前已经饱和;
  • 混合精度计算中出现溢出;
  • 梯度在归一化前已经产生异常;
  • 残差分支之间尺度严重失衡。

归一化并不能撤销已经发生的非线性饱和或浮点溢出。

2. 归一化的位置决定它能修正什么

考虑两个结构:

Linear -> ReLU -> LayerNorm

和:

LayerNorm -> Linear -> ReLU

第一种结构中,ReLU 的尺度变化发生在 LayerNorm 之前;第二种结构中,Linear 的输入先被标准化。二者的梯度路径、激活分布和残差行为不同。

Transformer 中常见两种结构:

Post-LN

x -> Attention -> Add(x) -> LayerNorm
x -> MLP       -> Add(x) -> LayerNorm

抽象形式:

y=LN(x+F(x))y=\operatorname{LN}(x+F(x))

Pre-LN

x -> LayerNorm -> Attention -> Add(x)
x -> LayerNorm -> MLP       -> Add(x)

抽象形式:

y=x+F(LN(x))y=x+F(\operatorname{LN}(x))

Pre-LN 把归一化放在子层输入处,通常能提供更直接的残差梯度路径;Post-LN 则在残差相加后归一化。二者不仅是代码顺序不同,还会改变初始化和深度扩展时的稳定性。

不能把“LayerNorm 适合 Transformer”理解为任意位置插入 LayerNorm 都等价。归一化位置必须与残差结构、注意力模块、MLP 模块和最终输出头一起分析。


十三、一个可运行的初始化对比实验

下面的程序比较 Xavier 和 Kaiming 在线性层加激活后的输出统计。它不用于证明某个初始化在所有任务上更好,而是展示初始化假设如何影响初始激活尺度。

import torch
from torch import nn

torch.manual_seed(42)

in_features = 512
out_features = 512
batch_size = 4096

x = torch.randn(batch_size, in_features)

def build_model(init_type: str) -> nn.Module:
    model = nn.Sequential(
        nn.Linear(in_features, out_features),
        nn.ReLU(),
    )

    linear = model[0]

    if init_type == "xavier":
        nn.init.xavier_normal_(linear.weight)
    elif init_type == "kaiming":
        nn.init.kaiming_normal_(
            linear.weight,
            mode="fan_in",
            nonlinearity="relu",
        )
    else:
        raise ValueError(f"unknown init_type: {init_type}")

    nn.init.zeros_(linear.bias)
    return model

for init_type in ("xavier", "kaiming"):
    model = build_model(init_type)

    with torch.no_grad():
        pre_activation = model[0](x)
        activation = model(x)

    print(f"\n{init_type}")
    print("weight std:", model[0].weight.std().item())
    print("pre-activation mean:", pre_activation.mean().item())
    print("pre-activation std:", pre_activation.std().item())
    print("ReLU output mean:", activation.mean().item())
    print("ReLU output std:", activation.std().item())
    print("ReLU zero ratio:", (activation == 0).float().mean().item())

在相同维度下,通常会观察到:

  • Kaiming 的权重标准差大于 Xavier;
  • Kaiming 经过 ReLU 后的二阶尺度更接近其设计目标;
  • ReLU 的零值比例在输入近似对称时通常接近一半;
  • Xavier 并不是“错误”,只是它没有专门补偿 ReLU 的半波截断。

输出中的具体数值会受随机种子、样本数量、PyTorch 版本和设备浮点实现影响,因此应关注统计趋势,而不是把某个固定数字作为规范保证。


十四、一个故意失败的反例:错误初始化导致深层信号衰减

下面构造一个不使用归一化、连续多个 ReLU 层的网络,并比较三种初始化:

  • 全部权重初始化为极小值;
  • Xavier;
  • Kaiming。
import torch
from torch import nn

torch.manual_seed(7)

class DeepMLP(nn.Module):
    def __init__(self, depth=20, width=256, init_type="kaiming"):
        super().__init__()

        layers = []
        for _ in range(depth):
            layers.append(nn.Linear(width, width))
            layers.append(nn.ReLU())

        self.net = nn.Sequential(*layers)

        for module in self.modules():
            if isinstance(module, nn.Linear):
                if init_type == "tiny":
                    nn.init.normal_(module.weight, mean=0.0, std=0.01)
                elif init_type == "xavier":
                    nn.init.xavier_normal_(module.weight)
                elif init_type == "kaiming":
                    nn.init.kaiming_normal_(
                        module.weight,
                        mode="fan_in",
                        nonlinearity="relu",
                    )
                else:
                    raise ValueError(init_type)

                nn.init.zeros_(module.bias)

    def forward(self, x):
        stats = []
        for module in self.net:
            x = module(x)
            if isinstance(module, nn.ReLU):
                stats.append({
                    "mean": x.mean().item(),
                    "std": x.std().item(),
                    "zero_ratio": (x == 0).float().mean().item(),
                })
        return x, stats

x = torch.randn(64, 256)

for init_type in ("tiny", "xavier", "kaiming"):
    model = DeepMLP(init_type=init_type)
    _, stats = model(x)

    print(f"\n{init_type}")
    print("first layer:", stats[0])
    print("last layer:", stats[-1])

tiny 初始化通常会表现为:

  • 第一层输出已经很小;
  • 后续层的输出继续衰减;
  • 最后几层激活接近全零;
  • 梯度也可能非常小。

Xavier 在 ReLU 网络中通常比极小初始化更合理,但它没有显式补偿 ReLU 的截断效应。Kaiming 则是针对该结构推导出的起点。

这仍不是“只要使用 Kaiming 就一定能训练”的证明。深度、残差、学习率、优化器、数据尺度、损失函数和归一化位置都会共同影响结果。


十五、如何诊断初始化或归一化问题

诊断时应区分“参数异常”“激活异常”“梯度异常”和“模式状态异常”。

1. 检查参数和激活统计

def summarize_parameters(model):
    for name, param in model.named_parameters():
        if param.requires_grad:
            print(
                name,
                "mean=", param.data.mean().item(),
                "std=", param.data.std().item(),
                "min=", param.data.min().item(),
                "max=", param.data.max().item(),
            )

def add_activation_hooks(model):
    hooks = []

    def make_hook(name):
        def hook(module, inputs, output):
            value = output.detach()
            print(
                name,
                "shape=", tuple(value.shape),
                "mean=", value.mean().item(),
                "std=", value.std().item(),
                "finite=", torch.isfinite(value).all().item(),
            )
        return hook

    for name, module in model.named_modules():
        if isinstance(module, (nn.Linear, nn.Conv2d, nn.BatchNorm1d,
                               nn.BatchNorm2d, nn.LayerNorm)):
            hooks.append(module.register_forward_hook(make_hook(name)))

    return hooks

使用 hook 后应在实验结束时移除:

hooks = add_activation_hooks(model)

try:
    _ = model(x)
finally:
    for hook in hooks:
        hook.remove()

否则重复运行诊断代码会注册越来越多的 hook,造成额外开销和重复输出。

2. 检查梯度

loss = model(x).pow(2).mean()
loss.backward()

for name, param in model.named_parameters():
    if param.grad is not None:
        print(
            name,
            "grad_norm=",
            param.grad.norm().item(),
            "finite=",
            torch.isfinite(param.grad).all().item(),
        )

典型现象包括:

  • 激活标准差逐层接近 0:可能是初始化过小、激活饱和或死 ReLU;
  • 激活标准差逐层快速增长:可能是初始化过大或残差分支尺度失衡;
  • 梯度范数接近 0:可能梯度消失、掩码过强或损失路径断开;
  • 梯度出现 NaNinf:可能学习率过大、混合精度溢出、输入异常或归一化数值不稳定;
  • 训练正常但评估异常:优先检查 BatchNorm 的 train/eval 状态和运行统计量。

3. 检查 BatchNorm 状态

for name, module in model.named_modules():
    if isinstance(module, (nn.BatchNorm1d, nn.BatchNorm2d)):
        print(name)
        print("training:", module.training)
        print("running_mean:", module.running_mean)
        print("running_var:", module.running_var)

如果 running_var 中出现极小值、异常大值或非有限值,应回溯:

  • 训练 batch 是否过小;
  • 输入是否存在极端异常值;
  • padding 是否参与统计;
  • 是否错误地在验证集上更新了 BatchNorm;
  • 是否恢复了完整的 state_dict

十六、常见错误与反例

错误一:看到 ReLU 就同时使用 Xavier 和 Kaiming

如果先调用 Xavier,又调用 Kaiming,后一次会覆盖前一次的权重。代码表面上可能看起来“都用了”,实际上只有最后一次初始化生效。

nn.init.xavier_uniform_(layer.weight)
nn.init.kaiming_normal_(layer.weight, nonlinearity="relu")

应根据实际激活函数选择一次初始化,并在代码中保持清晰。

错误二:使用了 Tanh,却按 ReLU 初始化

layer = nn.Linear(128, 128)
activation = nn.Tanh()

nn.init.kaiming_normal_(
    layer.weight,
    nonlinearity="relu",
)

这里初始化假设与实际激活不一致。ReLU 的半波截断假设不能直接套用到 Tanh。

错误三:把 LayerNorm 当作 BatchNorm

LayerNorm 不会利用 batch 中其他样本的统计量。增大 batch size 不会像 BatchNorm 那样改变 LayerNorm 的统计估计。

反过来,BatchNorm 也不是“对每个样本独立标准化”。两个样本放进同一个 batch,可能改变彼此的归一化结果。

错误四:验证和部署时忘记 eval()

model.eval()

with torch.no_grad():
    prediction = model(x)

torch.no_grad() 只关闭梯度记录,不会把 BatchNorm 切换到推理统计,也不会关闭 Dropout。两者职责不同。

错误五:在 LayerNorm 上设置错误的维度

对于 [N, T, D] 输入,nn.LayerNorm(D) 通常表示对每个 token 的隐藏维度归一化;nn.LayerNorm((T, D)) 则表示对整个序列和隐藏维度一起归一化。后者改变了统计语义,并且通常要求固定序列长度。

错误六:以为归一化后所有值都不超过某个范围

标准化通常控制均值和方差,不保证:

yi[1,1]y_i\in[-1,1]

经过 γ\gammaβ\beta 后,输出更不受固定区间约束。若业务需要边界约束,应使用专门的激活函数或投影方法,不能把归一化当作截断。


十七、混合精度和数值稳定性

归一化包含:

xμσ2+ϵ\frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}

当输入精度较低或方差很小时,ϵ\epsilon 会影响数值稳定性。它不是用来修复错误输入分布的万能参数,而是防止除零和极小分母。

在混合精度训练中,应关注:

  • 归一化实现是否使用适当的累加精度;
  • 输入是否出现 infNaN
  • eps 是否过小;
  • 梯度缩放是否正确;
  • 初始化是否使第一轮前向就产生极端值。

不要仅通过盲目增大 eps 来掩盖初始化、学习率或数据异常。增大 eps 会改变归一化结果,尤其在真实方差很小时影响明显。


十八、生成式 AI 和 Transformer 中的实际取舍

在 Transformer 中,LayerNorm 或其变体通常比 BatchNorm 更自然,原因来自数据流:

  1. 序列长度可能变化;
  2. 推理 batch 可能为 1;
  3. 自回归生成每次可能只处理一个或少量 token;
  4. 不希望一个请求的 token 统计影响另一个请求;
  5. BatchNorm 的运行统计量不适合频繁变化的生成上下文。

Transformer 的主要线性层一般不会简单地在每一层都使用传统 CNN 式 BatchNorm。模型是否使用 Pre-LN、Post-LN、RMSNorm 或其他归一化变体,需要结合注意力、MLP、残差和输出层共同分析。

对生成式模型还应保留以下状态:

  • 初始化后的参数;
  • 归一化层的可学习参数;
  • BatchNorm 若存在则包括运行统计量;
  • 优化器状态;
  • 混合精度训练的梯度缩放状态;
  • 与训练数据预处理一致的输入尺度配置。

这些状态共同决定恢复训练和部署推理的行为。只保存“看起来像权重的参数”可能不足以复现模型。


十九、从初始化到部署的验证流程

一个可靠的验证流程应按因果顺序执行:

第一步:确认张量语义

明确每个维度代表什么:

[N, C, H, W]      # 图像卷积
[N, T, D]         # Transformer 序列
[N, D]            # 普通特征

只有先知道维度语义,才能正确选择 BatchNorm 类型和 LayerNorm 的 normalized_shape

第二步:确认激活函数

记录每个线性层后实际使用的激活:

Linear -> Tanh
Linear -> ReLU
Linear -> GELU

初始化公式应基于实际激活,而不是类名或历史代码。

第三步:检查初始化后的单层统计

输入标准差约为 1 的随机张量,观察:

  • 预激活均值和标准差;
  • 激活后的均值和标准差;
  • ReLU 零值比例;
  • 是否存在非有限值。

第四步:检查深层传播

记录每一层的激活和梯度统计。如果尺度逐层呈指数级衰减或增长,优先检查:

  • fan_in 是否正确;
  • 初始化是否与激活匹配;
  • 残差分支是否需要缩放;
  • 归一化位置是否正确;
  • 输入是否已经经过重复或错误的缩放。

第五步:分别验证训练和推理

至少运行:

model.train()
train_output = model(x)

model.eval()
eval_output = model(x)

如果模型包含 BatchNorm,应确认差异是否来自预期的运行统计,而不是模式遗漏或状态未保存。

第六步:保存并恢复状态

torch.save(model.state_dict(), "model.pt")

restored = build_model()
restored.load_state_dict(torch.load("model.pt", map_location="cpu"))
restored.eval()

恢复后用同一输入比较输出。若差异明显,检查随机层、BatchNorm 状态、设备和 dtype 是否一致。


二十、核心判断

Xavier 和 Kaiming 都是在控制方差传播,但假设不同:

Xavier:Var(w)2fan_in+fan_out\text{Xavier:}\quad \operatorname{Var}(w) \approx \frac{2}{fan\_in+fan\_out}

Kaiming:Var(w)2fan_infor ReLU\text{Kaiming:}\quad \operatorname{Var}(w) \approx \frac{2}{fan\_in} \quad\text{for ReLU}

BatchNorm 和 LayerNorm 的关键差异不在于名字,而在于统计维度:

  • BatchNorm 主要按特征跨 batch 统计,并在训练和推理阶段使用不同统计来源;
  • LayerNorm 对每个样本内部的特征统计,通常不依赖 batch,也不需要运行均值和方差。

初始化解决“训练开始时信号处于什么尺度”,归一化解决“训练过程中激活尺度如何被重新控制”。二者必须与激活函数、权重布局、残差结构、训练模式、推理状态和模型保存方式一起设计。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。