共计 3151 个字符,预计需要花费 8 分钟才能阅读完成。
背景介绍
BatchNormalization(BN)是深度学习中的一项关键技术,由 Ioffe 和 Szegedy 在 2015 年提出。它的主要作用是解决深度神经网络训练过程中的 内部协变量偏移 问题,即网络中间层输入分布随着参数更新而不断变化,导致训练困难。BN 通过对每一层的输入进行标准化处理,使得数据分布更加稳定,从而允许使用更大的学习率,加速模型收敛,并具有一定的正则化效果。

数学推导
正向传播
给定一个 mini-batch 的输入数据 $X \in \mathbb{R}^{N \times D}$,其中 $N$ 是 batch size,$D$ 是特征维度,BN 的正向传播过程如下:
-
计算 batch 均值:
$$\mu = \frac{1}{N} \sum_{i=1}^N x_i$$ -
计算 batch 方差:
$$\sigma^2 = \frac{1}{N} \sum_{i=1}^N (x_i – \mu)^2$$ -
归一化处理:
$$\hat{x}_i = \frac{x_i – \mu}{\sqrt{\sigma^2 + \epsilon}}$$ -
缩放和平移(引入可学习参数 $\gamma$ 和 $\beta$):
$$y_i = \gamma \hat{x}_i + \beta$$
反向传播推导
反向传播需要计算损失 $L$ 对各个参数的梯度。我们采用链式法则逐步推导:
-
计算 $\frac{\partial L}{\partial y_i}$:
这是来自上一层的梯度,直接传递。 -
计算 $\frac{\partial L}{\partial \gamma}$ 和 $\frac{\partial L}{\partial \beta}$:
$$\frac{\partial L}{\partial \gamma} = \sum_{i=1}^N \frac{\partial L}{\partial y_i} \hat{x}i$$
$$\frac{\partial L}{\partial \beta} = \sum$$}^N \frac{\partial L}{\partial y_i -
计算 $\frac{\partial L}{\partial \hat{x}_i}$:
$$\frac{\partial L}{\partial \hat{x}_i} = \frac{\partial L}{\partial y_i} \gamma$$ -
计算 $\frac{\partial L}{\partial \sigma^2}$:
$$\frac{\partial L}{\partial \sigma^2} = \sum_{i=1}^N \frac{\partial L}{\partial \hat{x}_i} (x_i – \mu) \left(-\frac{1}{2} \right) (\sigma^2 + \epsilon)^{-3/2}$$ -
计算 $\frac{\partial L}{\partial \mu}$:
$$\frac{\partial L}{\partial \mu} = \left(\sum_{i=1}^N \frac{\partial L}{\partial \hat{x}i} \frac{-1}{\sqrt{\sigma^2 + \epsilon}} \right) + \frac{\partial L}{\partial \sigma^2} \frac{-2 \sum$$}^N (x_i – \mu)}{N -
最终计算 $\frac{\partial L}{\partial x_i}$:
$$\frac{\partial L}{\partial x_i} = \frac{\partial L}{\partial \hat{x}_i} \frac{1}{\sqrt{\sigma^2 + \epsilon}} + \frac{\partial L}{\partial \sigma^2} \frac{2(x_i – \mu)}{N} + \frac{\partial L}{\partial \mu} \frac{1}{N}$$
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.eps = eps
self.momentum = momentum
# 可训练参数
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))
def forward(self, x):
if self.training:
# 训练模式:使用当前 batch 统计量
mean = x.mean(dim=0)
var = x.var(dim=0, unbiased=False)
# 更新 running_mean 和 running_var
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:
# 推理模式:使用 running 统计量
mean = self.running_mean
var = self.running_var
# 归一化
x_hat = (x - mean) / torch.sqrt(var + self.eps)
# 缩放和平移
out = self.gamma * x_hat + self.beta
return out
训练与推理的差异
- 训练阶段:
- 使用当前 mini-batch 的均值和方差
-
更新 running_mean 和 running_var
-
推理阶段:
- 使用训练阶段累积的 running_mean 和 running_var
- 不再计算 batch 统计量
实验对比
我们构建一个简单的全连接网络,分别在 MNIST 数据集上测试使用 BN 和不使用 BN 的效果:
# 构造模型
model_with_bn = nn.Sequential(nn.Linear(784, 256),
BatchNorm1d(256),
nn.ReLU(),
nn.Linear(256, 10)
)
model_without_bn = nn.Sequential(nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10)
)
# 训练过程...
实验结果会显示:
– 使用 BN 的模型收敛更快
– 使用 BN 的模型对学习率的选择更鲁棒
– 使用 BN 的模型最终准确率通常更高
最佳实践
- 学习率设置:
- BN 允许使用更大的学习率
-
但过大学习率仍可能导致不稳定
-
batch size 影响:
- 较小的 batch size 会导致统计量估计不准确
-
建议 batch size 至少为 32
-
常见问题:
- 训练和测试模式切换错误:确保.eval()和.train()正确调用
- 初始化问题:gamma 初始化为 1,beta 初始化为 0
- 与其他正则化方法组合:与 dropout 一起使用时可能需要调整参数
总结
BatchNormalization 是深度学习中极其重要的技术,理解其反向传播过程有助于:
– 更深入地调试模型
– 自定义特殊归一化层
– 解决训练中的不稳定问题
虽然推导过程略显复杂,但掌握了基本原理后,可以灵活应用到各种场景中。建议读者在理解本文内容后,尝试手动实现 BN 层,并与标准实现对比,加深理解。
