共计 3187 个字符,预计需要花费 8 分钟才能阅读完成。
背景痛点
Batch Normalization(BN)在训练阶段需要维护运行时统计量(均值 / 方差),这使得其反向传播比普通层更复杂。关键在于:

- 梯度路径分叉 :损失
L对输入x_i的梯度需同时考虑x_i对归一化结果y_i的直接影响,以及通过全局均值μ和方差σ²的间接影响 - 链式法则嵌套:计算∂L/∂x_i 时需展开三层链式法则:
- 归一化输出
y_i = (x_i - μ)/√(σ² + ε) - 均值
μ = 1/m ∑x_j - 方差
σ² = 1/m ∑(x_j - μ)² - 计算图膨胀 :每个样本
x_i的梯度计算都依赖全体样本的统计量,导致显存访问模式复杂化
数学推导
设 batch 大小为 m,输入x_i 的梯度公式推导如下(原始论文 [1] 附录 A):
-
展开归一化变换:
$$y_i = \frac{x_i – μ}{\sqrt{σ^2 + ε}}$$ -
损失对
x_i的总梯度:
$$\frac{∂L}{∂x_i} = \frac{∂L}{∂y_i} \cdot \frac{∂y_i}{∂x_i} + \frac{∂L}{∂μ} \cdot \frac{∂μ}{∂x_i} + \frac{∂L}{∂σ^2} \cdot \frac{∂σ^2}{∂x_i}$$ -
逐项计算(关键步骤):
- 直接梯度项:
$$\frac{∂y_i}{∂x_i} = \frac{1}{\sqrt{σ^2 + ε}}$$ - 均值相关项:
$$\frac{∂μ}{∂x_i} = \frac{1}{m}$$
$$\frac{∂L}{∂μ} = \sum_{j=1}^m \frac{∂L}{∂y_j} \cdot \frac{-1}{\sqrt{σ^2 + ε}}$$ -
方差相关项:
$$\frac{∂σ^2}{∂x_i} = \frac{2(x_i – μ)}{m}$$
$$\frac{∂L}{∂σ^2} = \sum_{j=1}^m \frac{∂L}{∂y_j} \cdot (x_j – μ) \cdot \frac{-1}{2}(σ^2 + ε)^{-3/2}$$ -
最终合并形式:
$$\frac{∂L}{∂x_i} = \frac{1}{m\sqrt{σ^2 + ε}} \left[m\frac{∂L}{∂y_i} – \sum_{j=1}^m \frac{∂L}{∂y_j} – (x_i – μ) \cdot \sum_{j=1}^m \frac{∂L}{∂y_j}(x_j – μ) \cdot \frac{1}{σ^2 + ε} \right]$$
PyTorch 实现解析
以 torch.nn.BatchNorm2d._backward() 为例(v1.12 源码):
-
梯度预处理:
# 将上层梯度 (dL/dy) 与标准化系数相乘 grad_output = grad_output * self.weight.view(1, -1, 1, 1) -
均值梯度计算:
# 对应公式中的 sum(dL/dy_j)项 grad_mean = torch.sum(grad_output, dim=(0, 2, 3), keepdim=True) -
方差梯度计算:
# 计算点乘项 sum(dL/dy_j * (x_j - μ)) dot_p = torch.sum(grad_output * (input - self.running_mean.view(1, -1, 1, 1)), dim=(0, 2, 3), keepdim=True ) -
最终梯度合成:
# 对应完整梯度公式 grad_input = (grad_output - grad_mean / N - dot_p * (input - mean) / (var + self.eps) / N ) / torch.sqrt(var + self.eps)
手动实现示例
完整 BatchNorm 层实现(训练模式):
class MyBatchNorm2d:
def __init__(self, num_features, eps=1e-5):
self.gamma = torch.ones(num_features)
self.beta = torch.zeros(num_features)
self.eps = eps
self.running_mean = torch.zeros(num_features)
self.running_var = torch.ones(num_features)
def forward(self, x):
if self.training:
dims = (0, 2, 3)
mean = x.mean(dims, keepdim=True)
var = x.var(dims, unbiased=False, keepdim=True)
self.running_mean = 0.9 * self.running_mean + 0.1 * mean.squeeze()
self.running_var = 0.9 * self.running_var + 0.1 * var.squeeze()
else:
mean, var = self.running_mean, self.running_var
x_hat = (x - mean) / torch.sqrt(var + self.eps)
return self.gamma * x_hat + self.beta
def backward(self, grad_output, x):
N = x.shape[0] * x.shape[2] * x.shape[3]
mean = x.mean((0, 2, 3), keepdim=True)
var = x.var((0, 2, 3), unbiased=False, keepdim=True)
grad_gamma = (grad_output * x_hat).sum((0, 2, 3))
grad_beta = grad_output.sum((0, 2, 3))
dx_hat = grad_output * self.gamma.view(1, -1, 1, 1)
dvar = torch.sum(dx_hat * (x - mean) * -0.5 * (var + self.eps)**(-1.5),
(0, 2, 3))
dmean = torch.sum(dx_hat * (-1 / torch.sqrt(var + self.eps)),
(0, 2, 3)) \
+ dvar * torch.mean(-2 * (x - mean), (0, 2, 3))
grad_input = (dx_hat / torch.sqrt(var + self.eps) +
dvar.view(1, -1, 1, 1) * 2 * (x - mean) / N +
dmean.view(1, -1, 1, 1) / N)
return grad_input, grad_gamma, grad_beta
避坑指南
- 模式切换陷阱
- 训练 / 测试模式必须显式切换:
model.train()和model.eval() -
验证时忘记
eval()会导致使用 batch 统计量而非 running 统计量 -
同步问题
- 多 GPU 训练时需同步各卡的
running_mean/var(PyTorch 中设置sync_bn=True) -
梯度检查点(gradient checkpointing)可能破坏统计量计算
-
数值稳定性
- 方差计算应使用
unbiased=False(与原始论文一致) eps值不宜小于1e-5(FP16 训练建议eps=1e-3)
性能考量
- 计算开销
- 前向传播增加约 15% FLOPs(主要来自均值 / 方差计算)
-
反向传播因梯度分叉增加约 25% 内存访问
-
训练加速
- 允许使用 2~4 倍大的学习率
- 减少对参数初始化的敏感度
- 实际训练速度可提升 1.5~3 倍(取决于网络结构)
延伸思考
- 为什么 BatchNorm 在 NLP 任务(如 Transformer)中效果不如 CV 任务显著?
- 如何设计替代方案以解决 BatchNorm 在小 batch size(<8)时的性能下降问题?
- 在元学习(MAML)等需要二级导数的场景中,BatchNorm 的反向传播需要哪些特殊处理?
[1] Ioffe & Szegedy. “Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift”. ICML 2015.
