共计 3375 个字符,预计需要花费 9 分钟才能阅读完成。
前向传播的统计量计算
BatchNorm 的核心思想是对每个 mini-batch 进行标准化处理。给定输入 $X \in \mathbb{R}^{N \times C}$(N 为 batch size,C 为通道数),前向传播过程分为三步:

-
计算当前 batch 的均值和方差:
$$\mu = \frac{1}{N} \sum_{i=1}^N x_i$$
$$\sigma^2 = \frac{1}{N} \sum_{i=1}^N (x_i – \mu)^2$$ -
标准化处理(加上 epsilon 防止除零):
$$\hat{x}_i = \frac{x_i – \mu}{\sqrt{\sigma^2 + \epsilon}}$$ -
缩放和平移:
$$y_i = \gamma \hat{x}_i + \beta$$
这个过程在 PyTorch 中的实现非常直观:
mean = x.mean(dim=0)
var = x.var(dim=0, unbiased=False) # 注意这里使用有偏估计
x_hat = (x - mean) / torch.sqrt(var + eps)
out = weight * x_hat + bias # weight 即 γ,bias 即 β
反向传播的数学推导
反向传播需要计算损失 L 对各个参数的梯度。根据链式法则,我们需要先求 $\partial L/\partial \hat{x}$,然后逐步回推:
-
首先计算 $\partial L/\partial \hat{x}$(来自上层梯度):
$$\frac{\partial L}{\partial \hat{x}_i} = \frac{\partial L}{\partial y_i} \cdot \gamma$$ -
然后计算 $\partial L/\partial \sigma^2$(需要聚合 batch 维度):
$$\frac{\partial L}{\partial \sigma^2} = \sum_{i=1}^N \frac{\partial L}{\partial \hat{x}_i} \cdot (x_i – \mu) \cdot \left(-\frac{1}{2}\right) (\sigma^2 + \epsilon)^{-3/2}$$ -
接着计算 $\partial L/\partial \mu$(包含两条路径):
$$\frac{\partial L}{\partial \mu} = \left(\sum_{i=1}^N \frac{\partial L}{\partial \hat{x}i} \cdot \frac{-1}{\sqrt{\sigma^2 + \epsilon}}\right) + \frac{\partial L}{\partial \sigma^2} \cdot \frac{-2}{N} \sum^N (x_i – \mu)$$ -
最终得到输入梯度 $\partial L/\partial x$:
$$\frac{\partial L}{\partial x_i} = \frac{\partial L}{\partial \hat{x}_i} \cdot \frac{1}{\sqrt{\sigma^2 + \epsilon}} + \frac{\partial L}{\partial \sigma^2} \cdot \frac{2(x_i – \mu)}{N} + \frac{\partial L}{\partial \mu} \cdot \frac{1}{N}$$
PyTorch 实现解析
普通 BatchNorm 实现
PyTorch 的 BatchNorm1d 关键实现位于torch/nn/modules/batchnorm.py。前向传播时会更新 running_mean 和 running_var:
def forward(self, input):
# 训练模式
if self.training:
# 计算当前 batch 统计量
mean = input.mean([0, 2])
var = input.var([0, 2], unbiased=False)
# 更新 running 统计量(动量法)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
# 标准化和仿射变换
return F.batch_norm(input, mean, var, self.weight, self.bias, self.training, self.momentum, self.eps)
多卡同步实现
SyncBatchNorm 的关键区别在于统计量的跨卡同步。PyTorch 通过进程组通信实现:
def forward(self, input):
if self.training:
# 各卡先计算本地统计量
mean = input.mean([0, 2])
var = input.var([0, 2], unbiased=False)
count = torch.tensor(input.size(0), device=input.device)
# 跨卡同步(使用 all_reduce 聚合)combined = torch.cat([mean, var, count.unsqueeze(0)])
dist.all_reduce(combined, op=dist.ReduceOp.SUM, group=self.process_group)
# 计算全局统计量
combined = combined / self.world_size
mean, var, count = torch.split(combined, [self.num_features, self.num_features, 1])
# 更新 running 统计量
self.running_mean = ... # 同普通 BN
self.running_var = ...
性能优化技巧
数值稳定性
当 batch size 非常大时,方差计算可能溢出。改进方案:
def stable_var(x, dim):
mean = x.mean(dim, keepdim=True)
# 使用 Welford 算法
m = (x - mean).square().sum(dim)
return m / x.size(dim)
推理优化
在推理阶段,可以将卷积 +BN 融合为一个卷积操作。数学证明:
原始计算:
$$y = \gamma \cdot \frac{(W * x + b) – \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta$$
等效变换为:
$$y = (\frac{\gamma W}{\sqrt{\sigma^2 + \epsilon}}) * x + (\frac{\gamma (b – \mu)}{\sqrt{\sigma^2 + \epsilon}} + \beta)$$
PyTorch 中可以通过 torch.quantization.fuse_modules 自动完成这个优化。
避坑指南
- 多卡训练同步:
- 确保所有卡上的统计量同步后再更新 running_mean/var
-
使用 SyncBatchNorm 时注意设置正确的 process_group
-
Batch Size= 1 的处理:
- 方案 1:改用 InstanceNorm(即 BatchNorm2d with affine=False)
- 方案 2:使用累计统计量而非当前 batch
if batch_size == 1: mean = self.running_mean var = self.running_var
思考题
- 为什么 LayerNorm 不需要 running 统计量?
- LayerNorm 对每个样本独立计算统计量,不依赖 batch 维度
-
其统计量计算是确定性的,不存在训练 / 推理差异
-
Meta-learning 中的改造方法:
- 方案 1:在 inner loop 中冻结 BN 统计量
- 方案 2:使用 TaskNorm(跨 task 计算统计量)
- 方案 3:采用更灵活的 Normalization 方式如 FilterResponseNorm
总结
BatchNorm 的反向传播虽然复杂,但理解其数学本质后就能灵活应对各种变体。实际使用时要注意:
– 训练 / 推理模式区分
– 多卡同步的正确实现
– 边缘情况的兜底处理
希望本文的推导和代码分析能帮助你更自信地使用和定制 BatchNorm 层。
