共计 2604 个字符,预计需要花费 7 分钟才能阅读完成。
1. 为什么需要 Batch Normalization?
在深度神经网络训练过程中,每层输入的分布会随着参数更新不断变化,这种现象称为 Internal Covariate Shift(内部协变量偏移)。传统解决方法是对输入数据进行归一化(Normalization),但这只作用于网络最底层。BN 的核心思想是: 对每一层的输入都进行归一化,使数据分布稳定在合适范围内。

- 传统归一化的局限性:
- 仅对原始输入有效,无法解决深层网络的内部偏移
- 网络中间层的数据分布可能变得极端(如 sigmoid 激活函数的饱和区)
2. 数学原理详解
BN 操作分为两个阶段:
-
标准化(Normalization):
$$\mu_B = \frac{1}{m}\sum_{i=1}^m x_i$$
$$\sigma_B^2 = \frac{1}{m}\sum_{i=1}^m (x_i – \mu_B)^2$$
$$\hat{x}_i = \frac{x_i – \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}$$ -
仿射变换(Affine Transformation):
$$y_i = \gamma \hat{x}_i + \beta$$
其中:
– $\mu_B$, $\sigma_B$ 是当前 batch 的均值和方差
– $\gamma$ (scale)和 $\beta$ (shift)是可学习参数
– $\epsilon$ 是为数值稳定性添加的小常数(通常 1e-5)
3. PyTorch 实现
import torch
import torch.nn as nn
class BatchNorm1d(nn.Module):
def __init__(self, num_features, eps=1e-5, momentum=0.1):
super().__init__()
self.gamma = nn.Parameter(torch.ones(num_features))
self.beta = nn.Parameter(torch.zeros(num_features))
self.register_buffer("running_mean", torch.zeros(num_features))
self.register_buffer("running_var", torch.ones(num_features))
self.eps = eps
self.momentum = momentum
def forward(self, x):
if self.training:
# 训练模式:使用当前 batch 统计量
mean = x.mean(dim=0)
var = x.var(dim=0, unbiased=False)
# 更新移动平均
with torch.no_grad():
self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * mean
self.running_var = (1 - self.momentum) * self.running_var + self.momentum * var
else:
# 推理模式:使用保存的统计量
mean = self.running_mean
var = self.running_var
x_hat = (x - mean) / torch.sqrt(var + self.eps)
return self.gamma * x_hat + self.beta
4. TensorFlow 实现
import tensorflow as tf
class BatchNorm1d(tf.keras.layers.Layer):
def __init__(self, num_features, eps=1e-5, momentum=0.1):
super().__init__()
self.gamma = tf.Variable(tf.ones(num_features))
self.beta = tf.Variable(tf.zeros(num_features))
self.moving_mean = tf.Variable(tf.zeros(num_features), trainable=False)
self.moving_var = tf.Variable(tf.ones(num_features), trainable=False)
self.eps = eps
self.momentum = momentum
def call(self, x, training=None):
if training:
mean, var = tf.nn.moments(x, axes=[0])
self.moving_mean.assign((1 - self.momentum) * self.moving_mean + self.momentum * mean)
self.moving_var.assign((1 - self.momentum) * self.moving_var + self.momentum * var)
else:
mean = self.moving_mean
var = self.moving_var
x_hat = (x - mean) / tf.sqrt(var + self.eps)
return self.gamma * x_hat + self.beta
5. 实验对比(CIFAR-10)
我们对比了 ResNet18 在 CIFAR-10 上使用 BN 前后的表现:
| 指标 | 无 BN | 有 BN |
|---|---|---|
| 训练损失收敛步数 | 15k | 8k |
| 测试集准确率 | 78.2% | 85.6% |
- 收敛速度:BN 使训练损失更快下降
- 模型性能:测试准确率提升 7% 以上
6. 生产环境建议
Batch Size 较小时的替代方案
当 batch size 小于 16 时,BN 的统计量可能不准确,推荐使用:
- Layer Normalization:对单个样本的所有特征做归一化
- Group Normalization:将通道分组后做归一化
模型导出注意事项
- 冻结 BN 层的
running_mean和running_var - 确保推理时
training=False - 检查移动平均统计量的初始化值
实现要点总结
- 训练 / 推理模式分离:两种模式下使用不同的统计量
- 移动平均更新:使用 momentum 平滑历史统计量
- 数值稳定性:添加 epsilon 防止除零错误
- 可学习参数:γ 和 β 保留模型的表达能力
通过这个实现,我们不仅理解了 BN 的数学本质,也掌握了其在生产环境中的应用技巧。建议读者在自己的数据集上尝试不同 normalization 方法的组合效果。
