AI 工程基础体系 · 第 62/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。
神经网络初始化与归一化:Xavier、Kaiming、BatchNorm 和 LayerNorm
神经网络训练开始时,参数通常是随机初始化的。初始化过小,信号和梯度会在多层传播后逐渐消失;初始化过大,激活值或梯度可能爆炸。归一化则在训练过程中重新调整中间表示的尺度,使优化问题更容易处理。
初始化和归一化解决的是相关但不同的问题:
- 初始化决定训练开始时参数和激活值的统计尺度。
- 归一化在前向传播中根据当前数据重新调整激活值。
- Xavier 初始化主要在输入输出尺度之间折中,常用于线性层、
tanh等近似对称激活函数。 - Kaiming 初始化针对 ReLU 类激活函数补偿其截断负半轴造成的方差损失。
- BatchNorm按批次统计特征,训练和推理阶段使用不同的统计来源。
- LayerNorm按单个样本的特征维度统计,训练和推理通常使用相同的计算方式。
理解这些方法,需要先从信号如何穿过网络开始。
一、为什么初始化会影响训练稳定性
设某一层是线性变换:
其中:
- 是输入;
- 是权重;
- 是偏置;
- 是输入维度,也称为
fan_in; - 是激活函数之前的预激活值。
先作几个常见近似:
- 输入和权重均值为 0;
- 权重和输入相互独立;
- 不同输入维度之间近似独立;
- 偏置初始为 0。
于是:
在独立条件下:
因此:
如果希望每层输出的方差大致保持不变,即:
就需要:
这只是线性层的前向传播条件。实际网络还包含非线性激活,而且反向传播还要求梯度尺度不要快速变大或变小。
1. 深层网络中的方差连乘
假设每一层都使信号方差乘以常数 ,经过 层后:
当 时,方差指数衰减;当 时,方差指数增长。即使 ,经过 50 层后也只有:
这就是“每层只缩小一点”仍然会导致深层信号消失的原因。
梯度也有类似问题。若每层的雅可比矩阵都使梯度范数平均乘以 ,则:
初始化方法的目标不是让每个样本、每个神经元的值完全相同,而是让整体统计尺度在训练初期处于合理范围。
二、激活函数会改变方差传播
初始化不能脱离激活函数讨论。设:
其中 是激活函数。
1. ReLU 会丢弃一半负值
ReLU 定义为:
如果 服从均值为 0 的对称分布,大约一半值为负,会被置为 0。对于零均值高斯变量 ,有:
因此,从“二阶矩”角度看,ReLU 使尺度大约减少一半。
需要注意,ReLU 输出的均值通常大于 0,所以它的方差并不严格等于输入方差的一半。若 ,则:
工程推导中常用“二阶矩约减半”的近似,因为它足以导出 Kaiming 初始化的主要尺度。
2. Sigmoid 和 Tanh 还可能进入饱和区
Sigmoid 为:
其导数为:
当 很大或很小时,Sigmoid 接近 1 或 0,导数接近 0,反向梯度会消失。
Tanh 的导数为:
当 较大时,Tanh 也会饱和。Xavier 初始化的一个重要目标,就是让线性输出的尺度不要过大,以便更多值处于激活函数的有效梯度区域。
三、Xavier 初始化:在前向和反向之间折中
Xavier 初始化也称为 Glorot 初始化。其核心思想是同时考虑:
- 前向传播中输入维度
fan_in; - 反向传播中输出维度
fan_out。
设某层有:
- 输入维度 ;
- 输出维度 。
前向传播若要保持方差,大致需要:
反向传播若要保持梯度方差,大致需要:
两者折中得到:
这就是 Xavier 正态初始化的常见形式:
若使用均匀分布:
均匀分布的方差为 ,令其等于上面的目标方差:
得到:
因此 Xavier 均匀初始化为:
1. 一个完整计算例子
考虑一个线性层:
nn.Linear(4, 2)
因此:
Xavier 初始化的目标方差为:
均匀分布边界为:
所以:
- Xavier 正态初始化:标准差为 ;
- Xavier 均匀初始化:从 中采样。
如果输入每个维度的方差约为 1,那么线性层输出的方差近似为:
这并不等于 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,前面已经得到:
线性层有:
为了让 ReLU 后的二阶矩大致保持不变,需要:
因此:
Kaiming 正态初始化为:
对应标准差:
Kaiming 均匀初始化的边界为:
因为均匀分布 的方差是 。
1. 继续使用前面的层作为例子
对于 nn.Linear(4, 2):
Kaiming 正态初始化目标方差为:
标准差为:
如果输入方差为 1,则线性输出的二阶矩近似为:
经过 ReLU 后:
这就是 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 定义为:
其中 是负半轴斜率。由于负半轴不再完全置零,所需的初始化增益不同。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_in 和 fan_out 如何计算
对于二维线性层权重:
nn.Linear(in_features, out_features)
权重形状通常为:
[out_features, in_features]
因此:
卷积层还要乘以卷积核空间尺寸。对于二维卷积权重形状:
[out_channels, in_channels, kernel_height, kernel_width]
有:
例如:
nn.Conv2d(
in_channels=3,
out_channels=64,
kernel_size=3,
)
其:
PyTorch 的初始化函数会根据参数张量形状计算这些值,但它对张量布局有约定。若自定义权重不是常见的线性层或卷积层布局,应确认 fan_in 和 fan_out 是否被正确解释。错误的布局会导致初始化尺度偏离预期。
一个常见反例是把矩阵转置后直接初始化:
weight = torch.empty(128, 64)
# 后续实际计算可能是 x @ weight,而不是 F.linear(x, weight)
如果实际计算语义与 PyTorch 默认 Linear 的权重布局不同,却仍按默认布局理解 fan_in,初始化尺度可能被交换。自定义层应明确写出:
再决定哪个维度是输入维度。
六、Xavier 与 Kaiming 的选择边界
可以用以下因果关系理解二者:
| 激活函数或结构 | 常见初始化起点 | 原因 |
|---|---|---|
| Linear | Xavier 或较简单的方差保持初始化 | 不存在截断效应 |
| Tanh | Xavier,并考虑 gain |
需要控制进入饱和区的概率 |
| ReLU | Kaiming,nonlinearity="relu" |
ReLU 约丢弃一半二阶矩 |
| Leaky ReLU | Kaiming,并指定负斜率 | 负半轴仍保留部分信号 |
| GELU、SiLU | 没有一个完全等价的简单公式 | 激活函数不是硬半波截断,通常结合架构和实验验证 |
最后一行很重要:不能把所有非线性函数都简单归入“ReLU,所以使用 Kaiming”。GELU 和 SiLU 的输入输出统计与 ReLU 不同。Transformer 中常见的线性层,通常还会配合 LayerNorm、残差连接和特定的缩放策略,因此不能只靠一个初始化公式推断整个模型的稳定性。
七、归一化到底在做什么
归一化层通常先计算某个维度集合上的均值和方差:
然后标准化:
最后应用可学习的仿射变换:
其中:
- 是防止除零的小常数;
- 是可学习的缩放参数;
- 是可学习的偏移参数。
归一化并不是简单地“把数据压缩到 0 到 1”。它通常使指定维度上的均值接近 0、方差接近 1,然后通过 允许网络恢复所需的尺度和偏移。
初始化与归一化的差异是:
- 初始化只在参数创建时发生一次;
- 归一化在每次前向传播时作用于激活;
- 归一化的统计维度决定它改变了哪些样本之间的关系;
- 归一化参数本身也会进入优化和模型检查点。
八、BatchNorm:按批次统计特征
BatchNorm 的典型输入形状取决于层类型。
对于 BatchNorm1d,常见输入为:
[N, C]
或:
[N, C, L]
其中:
- 是 batch size;
- 是通道或特征数;
- 是额外的序列长度。
BatchNorm 通常对每个通道 独立统计,统计维度包括 batch 维以及可能的空间或序列维度。对于二维图像输入 [N, C, H, W],BatchNorm2d 对每个通道统计 这些维度。
对于固定通道 ,训练时可以抽象为:
这里的 是当前 BatchNorm 实际统计到的元素数量。
1. BatchNorm 的训练状态和推理状态
BatchNorm 有两套统计来源:
训练模式
model.train()
通常使用当前 mini-batch 的均值和方差,并更新运行统计量:
running_meanrunning_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,因此存在几个真实边界:
- batch 太小:均值和方差估计噪声大。
- 分布式训练:每张 GPU 只看到本地 batch,局部统计可能与全局统计不同。
- 数据分布变化:运行统计量来自历史训练数据,部署数据发生漂移时可能失配。
- 变长序列和 padding:若 padding 位置参与统计,统计量可能被无效位置污染。
- 训练和推理模式错误:忘记
eval()会使推理输出受当前 batch 影响。 - 状态保存不完整:只保存参数而不保存 BatchNorm 的运行统计量,恢复后行为可能改变。
在多卡训练中,如果模型确实需要跨设备共享 batch 统计,通常要考虑同步 BatchNorm 等机制;但同步会引入跨设备通信成本,也可能降低吞吐。是否使用它取决于 batch 大小、模型结构和训练系统,而不是仅凭名称选择。
九、LayerNorm:按单个样本的特征维度统计
LayerNorm 不依赖 batch 中的其他样本。对输入最后若干个维度计算均值和方差。
对于形状:
[N, T, D]
如果使用:
nn.LayerNorm(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 风格的统计
如果把每个特征 作为一个通道,BatchNorm 可能在 维度上统计:
它回答的是:
在整个 batch 和序列位置中,第 个特征的总体均值是多少?
LayerNorm 的统计
LayerNorm 对每个 单独统计:
它回答的是:
对这个样本的这个 token,所有 hidden features 的均值是多少?
因此两者不是同一种归一化的不同 API 名称,而是统计对象不同:
| 特性 | BatchNorm | LayerNorm |
|---|---|---|
| 统计对象 | 同一特征跨样本,可能还跨空间或时间 | 单个样本内部的特征 |
| 是否依赖 batch | 是 | 否 |
| 是否有运行统计量 | 通常有 | 没有 BatchNorm 式运行统计量 |
| batch size 为 1 的稳定性 | 可能较差 | 通常不受影响 |
| 常见场景 | CNN、视觉模型 | Transformer、序列模型 |
| 推理行为 | 依赖 train/eval 状态 |
通常训练推理计算方式一致 |
十一、归一化层中的可学习参数与初始化
归一化层通常包含:
weight,对应 ,初始为 1;bias,对应 ,初始为 0。
这样初始化后,归一化层一开始主要执行标准化,不额外改变尺度和均值:
例如:
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_mean 和 running_var 不是普通的梯度参数,通常不通过反向传播更新,而是作为模块状态在训练时更新。使用 state_dict() 保存模型时,应同时保存这些状态。
state = model.state_dict()
torch.save(state, "checkpoint.pt")
如果只手动保存 named_parameters() 得到的可训练参数,却遗漏 BatchNorm 的运行统计量,恢复后的推理结果可能与保存前不一致。
十二、初始化与归一化如何共同工作
一个常见误解是:
既然有 BatchNorm 或 LayerNorm,初始化就不重要了。
归一化会减弱前几层尺度错误的影响,但不能消除所有初始化问题。
1. 归一化之前的数值仍可能溢出
如果权重初始化过大,线性层可能先产生极大的 z:
之后即使进入归一化层,也可能出现:
- 激活函数在归一化前已经饱和;
- 混合精度计算中出现溢出;
- 梯度在归一化前已经产生异常;
- 残差分支之间尺度严重失衡。
归一化并不能撤销已经发生的非线性饱和或浮点溢出。
2. 归一化的位置决定它能修正什么
考虑两个结构:
Linear -> ReLU -> LayerNorm
和:
LayerNorm -> Linear -> ReLU
第一种结构中,ReLU 的尺度变化发生在 LayerNorm 之前;第二种结构中,Linear 的输入先被标准化。二者的梯度路径、激活分布和残差行为不同。
Transformer 中常见两种结构:
Post-LN
x -> Attention -> Add(x) -> LayerNorm
x -> MLP -> Add(x) -> LayerNorm
抽象形式:
Pre-LN
x -> LayerNorm -> Attention -> Add(x)
x -> LayerNorm -> MLP -> Add(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:可能梯度消失、掩码过强或损失路径断开;
- 梯度出现
NaN或inf:可能学习率过大、混合精度溢出、输入异常或归一化数值不稳定; - 训练正常但评估异常:优先检查 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)) 则表示对整个序列和隐藏维度一起归一化。后者改变了统计语义,并且通常要求固定序列长度。
错误六:以为归一化后所有值都不超过某个范围
标准化通常控制均值和方差,不保证:
经过 和 后,输出更不受固定区间约束。若业务需要边界约束,应使用专门的激活函数或投影方法,不能把归一化当作截断。
十七、混合精度和数值稳定性
归一化包含:
当输入精度较低或方差很小时, 会影响数值稳定性。它不是用来修复错误输入分布的万能参数,而是防止除零和极小分母。
在混合精度训练中,应关注:
- 归一化实现是否使用适当的累加精度;
- 输入是否出现
inf或NaN; eps是否过小;- 梯度缩放是否正确;
- 初始化是否使第一轮前向就产生极端值。
不要仅通过盲目增大 eps 来掩盖初始化、学习率或数据异常。增大 eps 会改变归一化结果,尤其在真实方差很小时影响明显。
十八、生成式 AI 和 Transformer 中的实际取舍
在 Transformer 中,LayerNorm 或其变体通常比 BatchNorm 更自然,原因来自数据流:
- 序列长度可能变化;
- 推理 batch 可能为 1;
- 自回归生成每次可能只处理一个或少量 token;
- 不希望一个请求的 token 统计影响另一个请求;
- 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 都是在控制方差传播,但假设不同:
BatchNorm 和 LayerNorm 的关键差异不在于名字,而在于统计维度:
- BatchNorm 主要按特征跨 batch 统计,并在训练和推理阶段使用不同统计来源;
- LayerNorm 对每个样本内部的特征统计,通常不依赖 batch,也不需要运行均值和方差。
初始化解决“训练开始时信号处于什么尺度”,归一化解决“训练过程中激活尺度如何被重新控制”。二者必须与激活函数、权重布局、残差结构、训练模式、推理状态和模型保存方式一起设计。
系列导航与关联阅读
- 系列入口:AI 工程完整学习路线:从机器学习与 Transformer 到 RAG、Agent 和生产治理
- 上一篇:深度学习损失与正则化:目标设计、Dropout、权重衰减和早停
- 下一篇:深度学习数据管道:Dataset、采样、增强、预取与吞吐诊断
官方资料
本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。

评论
0 条讨论