共计 1566 个字符,预计需要花费 4 分钟才能阅读完成。
背景与痛点
深度神经网络训练过程中,每一层的输入分布会随着前一层参数更新而不断变化,这种现象被称为内部协变量偏移(Internal Covariate Shift)。这会导致后续层需要不断适应新的数据分布,从而降低训练效率。Batch Norm 通过规范化每一层的输入分布,有效缓解了这一问题。

- 内部协变量偏移的影响:随着网络深度增加,层间分布变化会累积放大,导致梯度消失或爆炸
- Batch Norm 的解决方案:对每个 batch 的数据进行标准化(减去均值、除以标准差),使输入保持稳定分布
- 额外收益:允许使用更大的学习率,减少对参数初始化的依赖,并有一定正则化效果
技术对比
不同的归一化技术适用于不同场景:
- Batch Norm:依赖 batch 统计量,在 CNN 中表现优异,但对 batch size 敏感
- Layer Norm:沿特征维度归一化,适合 RNN 和 Transformer 结构
- Instance Norm:对每个样本单独归一化,常用于风格迁移任务
数学原理
前向传播
对于 batch 中的特征 x:
- 计算 batch 均值:
μ = mean(x) - 计算 batch 方差:
σ² = var(x) - 归一化:
x̂ = (x - μ)/√(σ² + ε) - 缩放和平移:
y = γx̂ + β
反向传播
在训练时,Batch Norm 需要维护移动平均的均值和方差用于推理阶段:
- 移动平均均值:
μ_running = momentum * μ_running + (1 - momentum) * μ - 移动平均方差:
σ²_running = momentum * σ²_running + (1 - momentum) * σ²
PyTorch 实现
import torch
import torch.nn as nn
class ModelWithBN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 64, kernel_size=3)
self.bn1 = nn.BatchNorm2d(64)
self.conv2 = nn.Conv2d(64, 128, kernel_size=3)
self.bn2 = nn.BatchNorm2d(128)
self.fc = nn.Linear(128*10*10, 10)
def forward(self, x):
x = self.conv1(x)
x = self.bn1(x)
x = F.relu(x)
x = self.conv2(x)
x = self.bn2(x)
x = F.relu(x)
x = x.view(x.size(0), -1)
x = self.fc(x)
return x
# 训练时自动计算 batch 统计量
model.train()
output = model(input)
# 推理时使用训练阶段累积的统计量
model.eval()
with torch.no_grad():
output = model(input)
调参指南
- momentum 选择:
- 典型值 0.1-0.3 用于小数据集
- 0.9-0.99 适合大数据集
-
影响移动平均的平滑程度
-
batch size 影响:
- 小 batch size 会导致统计量估计不准确
- 建议 batch size≥32
-
可考虑 Group Norm 替代
-
初始 γ 和 β :
- γ 初始化为 1,β 初始化为 0
- 允许网络学习是否使用归一化
避坑实践
- 混合精度训练:
- 需设置
model.bn.float()保持 BN 层为 FP32 -
避免数值下溢
-
同步 BN:
- 多 GPU 训练时考虑 SyncBatchNorm
-
确保跨卡统计量同步
-
权重衰减:
- 通常不对 BN 的 γ 和 β 使用权重衰减
- 避免过度约束缩放参数
开放性问题
在小 batch size 场景下,Batch Norm 的统计量估计会变得不可靠。可能的改进方向包括:
- 使用 Group Norm 或 Layer Norm 替代
- 跨 batch 累积统计量
- 预计算数据集全局统计量
- 开发更鲁棒的小 batch 归一化方法
这些方法各有优劣,需要根据具体任务进行选择和调整。
正文完
发表至: 深度学习
近三天内
