共计 2494 个字符,预计需要花费 7 分钟才能阅读完成。
为什么 BatchNorm 需要特殊处理?
BatchNorm 在训练和推理阶段的行为差异是第一个需要理解的要点。训练时,它用当前 batch 的均值 / 方差做归一化:

$$\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$$
而在 eval 模式时,则使用全局统计量 running_mean/running_var。这种差异导致手动实现时容易忽略对 $\mu_B$ 和 $\sigma_B^2$ 的梯度计算——它们是前向传播时的中间变量,但反向传播时也需要参与链式法则。
数学推导:拆解梯度计算
前向传播步骤
- 计算 batch 均值 $\mu_B$
- 计算 batch 方差 $\sigma_B^2$
- 归一化:$\hat{x}_i = \frac{x_i-\mu_B}{\sqrt{\sigma_B^2+\epsilon}}$
- 缩放平移:$y_i = \gamma\hat{x}_i + \beta$
反向传播关键点
对输入 $x_i$ 的梯度包含三部分:
$$\frac{\partial L}{\partial x_i} = \frac{\partial L}{\partial y_i} \cdot \gamma \cdot \left(\frac{1}{\sqrt{\sigma_B^2+\epsilon}} – \frac{(x_i-\mu_B)^2}{m(\sigma_B^2+\epsilon)^{3/2}}\right) + \frac{\partial L}{\partial \mu_B}\cdot\frac{-1}{m} + \frac{\partial L}{\partial \sigma_B^2}\cdot\frac{-2(x_i-\mu_B)}{m}$$
其中对 $\mu_B$ 和 $\sigma_B^2$ 的梯度常被忽略,这是手动实现的主要难点。完整推导建议参考 原论文 的补充材料。
PyTorch 实现对比
手动实现版
import torch
import torch.nn as nn
class ManualBatchNorm2d(nn.Module):
def __init__(self, num_features, eps=1e-5, momentum=0.1):
super().__init__()
self.gamma = nn.Parameter(torch.ones(num_features, 1, 1))
self.beta = nn.Parameter(torch.zeros(num_features, 1, 1))
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, 2, 3), keepdim=True)
var = x.var(dim=(0, 2, 3), unbiased=False, keepdim=True)
# 更新 running 统计量
self.running_mean = (1-self.momentum)*self.running_mean + self.momentum*mean.squeeze()
self.running_var = (1-self.momentum)*self.running_var + self.momentum*var.squeeze()
# 归一化
x_hat = (x - mean) / torch.sqrt(var + self.eps)
else:
x_hat = (x - self.running_mean.view(1,-1,1,1)) / \
torch.sqrt(self.running_var.view(1,-1,1,1) + self.eps)
return self.gamma * x_hat + self.beta
自动求导版
class AutoBatchNorm2d(nn.Module):
def __init__(self, num_features):
super().__init__()
self.bn = nn.BatchNorm2d(num_features, affine=False)
self.gamma = nn.Parameter(torch.ones(num_features, 1, 1))
self.beta = nn.Parameter(torch.zeros(num_features, 1, 1))
def forward(self, x):
return self.gamma * self.bn(x) + self.beta
四大避坑指南
- 模式切换陷阱:
- eval 模式必须冻结 running_mean/running_var
-
测试时若忘记调用
model.eval()会导致统计量漂移 -
小 batch_size 问题:
- 当 batch_size= 1 时方差计算为零
-
解决方案:使用
SyncBatchNorm或改用 LayerNorm -
分布式训练同步:
- 多 GPU 时需同步各卡的统计量
-
PyTorch 的
nn.BatchNorm2d已内置处理 -
数值稳定性:
- 分母添加 epsilon(典型值 1e-5)
- 梯度爆炸时可尝试调整 momentum 值
实验验证
在 CIFAR-10 上对比三种实现:
- 手动实现 BatchNorm
- 自动求导版本
- 原生
nn.BatchNorm2d
训练曲线显示三者在验证准确率上差异不超过 0.5%,但手动实现的训练时间增加约 15%。梯度分布可视化表明原生实现数值稳定性更优。
延伸思考
- 为什么 LayerNorm 不需要 running 统计量?
-
LayerNorm 对单个样本做归一化,不依赖 batch 维度
-
如何实现 per-channel 的 BatchNorm?
- 修改 gamma/beta 的形状为
(num_features,) - 统计量计算保持 channel 维度
完整代码见GitHub 示例
