共计 2733 个字符,预计需要花费 7 分钟才能阅读完成。
背景与作用
BatchNormalization(BN)是深度学习中广泛使用的技术,主要用于解决内部协变量偏移(Internal Covariate Shift)问题。简单来说,随着网络层数的加深,每层输入的分布会逐渐发生变化,导致训练过程变慢。BN 通过对每一层的输入进行标准化处理,使得输入分布保持稳定,从而加速训练过程并提高模型性能。

数学推导
前向传播回顾
BN 的前向传播过程主要包括以下步骤:
-
计算 batch 的均值:
$$\mu_B = \frac{1}{m} \sum_{i=1}^m x_i$$ -
计算 batch 的方差:
$$\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}}$$ -
缩放和平移:
$$y_i = \gamma \hat{x}_i + \beta$$
反向传播推导
反向传播的核心是计算损失函数对输入 $x_i$、缩放参数 $\gamma$ 和偏移参数 $\beta$ 的梯度。以下是详细的推导过程:
-
损失函数对 $y_i$ 的梯度:
$$\frac{\partial L}{\partial y_i}$$ -
对 $\beta$ 的梯度:
$$\frac{\partial L}{\partial \beta} = \sum_{i=1}^m \frac{\partial L}{\partial y_i}$$ -
对 $\gamma$ 的梯度:
$$\frac{\partial L}{\partial \gamma} = \sum_{i=1}^m \frac{\partial L}{\partial y_i} \hat{x}_i$$ -
对 $\hat{x}_i$ 的梯度:
$$\frac{\partial L}{\partial \hat{x}_i} = \frac{\partial L}{\partial y_i} \gamma$$ -
对 $\sigma_B^2$ 的梯度:
$$\frac{\partial L}{\partial \sigma_B^2} = \sum_{i=1}^m \frac{\partial L}{\partial \hat{x}_i} (x_i – \mu_B) \left(-\frac{1}{2} (\sigma_B^2 + \epsilon)^{-3/2}\right)$$ -
对 $\mu_B$ 的梯度:
$$\frac{\partial L}{\partial \mu_B} = \left(\sum_{i=1}^m \frac{\partial L}{\partial \hat{x}i} \frac{-1}{\sqrt{\sigma_B^2 + \epsilon}}\right) + \frac{\partial L}{\partial \sigma_B^2} \frac{-2}{m} \sum^m (x_i – \mu_B)$$ -
对 $x_i$ 的梯度:
$$\frac{\partial L}{\partial x_i} = \frac{\partial L}{\partial \hat{x}_i} \frac{1}{\sqrt{\sigma_B^2 + \epsilon}} + \frac{\partial L}{\partial \sigma_B^2} \frac{2(x_i – \mu_B)}{m} + \frac{\partial L}{\partial \mu_B} \frac{1}{m}$$
代码实现
以下是 PyTorch 中自定义 BN 层的完整实现:
import torch
import torch.nn as nn
class BatchNorm1dCustom(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.eps = eps
self.momentum = momentum
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:
mean = x.mean(dim=0)
var = x.var(dim=0, unbiased=False)
# Update running statistics
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)
out = self.gamma * x_hat + self.beta
return out
避坑指南
数值稳定性问题
- 训练初期方差接近 0 时,可能导致除零错误。解决方案:
- 添加一个小的常数 $\epsilon$(通常 1e-5)
-
使用更稳定的计算方式,如
torch.var时设置unbiased=False -
Batch size 较小时的影响:
- 小 batch size 会导致统计量估计不准确
-
解决方案:
- 使用更大的 batch size
- 采用 Group Normalization 等替代方法
-
推理模式注意事项:
- 务必使用训练阶段计算的 running_mean 和 running_var
- 确保模型在 eval()模式下运行
性能考量
- 训练速度:
- BN 能显著加快训练收敛速度
-
但每个 batch 需要额外的计算开销
-
内存占用:
- 需要存储 running_mean 和 running_var
- 对于大模型,可能增加显存压力
思考与扩展
- 不同初始化方法对 BN 效果的影响:
- 权重初始化应与 BN 配合使用
-
例如,使用 He 初始化时,可以适当增大学习率
-
尝试其他 Normalization 方法:
- Layer Normalization
- Instance Normalization
- Group Normalization
- 比较它们在特定任务上的效果
通过本文的学习,你应该对 BN 的反向传播原理有了深入理解,并能够在实际项目中正确实现和应用 BN 层。接下来可以尝试在不同网络结构中应用 BN,并观察其对模型性能的影响。
