ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Batch Normalization原理与PyTorch实战:从Internal Covariate Shift到稳定训练

Batch Normalization原理与PyTorch实战:从Internal Covariate Shift到稳定训练 真正开始训深层网络之后很多人都会遇到一个让人非常头疼的场景网络层数一加深loss像被定住一样不降或者一上来直接NaN就算勉强降了训练过程也是忽上忽下换个初始化方式结果又完全不一样。这个问题的“罪魁祸首”之一就是深度学习中非常经典的 Internal Covariate Shift 问题而 Batch Normalization 的出现可以说把整个深度网络的训练体验拔高了一大截。这篇博文就围绕这个主题把 Internal Covariate Shift 到底是什么、Batch Normalization 为什么管用、以及实际用 PyTorch 落地时有哪些细节和坑尽量一次性讲透。适合刚入门深度学习、正在调 CNN 模型的同学也适合那些训练经常不稳定、想搞清楚 BN 背后原理的工程师。1. 先弄清 Internal Covariate Shift 到底是什么1.1 网络内部的数据分布其实一直在“地震”要理解 Internal Covariate Shift先记住一个事实深度网络是一层层堆起来的每一层都在学习一个映射函数。假设输入 x 经过第一层得到 h1h1 再经过第二层得到 h2以此类推。这里的关键在于第二层看到的数据是 h1而不是原始的 x。问题就出在这里。训练过程中第一层的参数在不停更新所以它输出的 h1 的分布也在不停变化。也就是说第二层今天的输入分布和明天的输入分布可能完全不是一回事。对于更深层的网络来说这种输入分布的变化会被逐层放大越靠后的层看到的输入分布越不稳定。这就像流水线上前面工位一直在换零件规格后面工位永远在适应新东西生产效率自然大打折扣。这里说的“数据分布变化”有个专门的数学刻画网络内部层的输入分布随着前面层参数更新而发生改变的现象就是 Internal Covariate Shift简称 ICS。它和传统机器学习里说的 Covariate Shift 还不完全一样。传统场景中训练集和测试集的分布不同这是外部环境变化导致的而 ICS 是发生在网络内部、由训练过程自身引起的分布漂移。1.2 为什么分布漂移会让训练变慢、变难ICS 对训练的负面影响主要体现在三个层面上。第一后层需要不断学习新分布。每一轮的输入分布都在变后层网络的参数必须反复调整去适应新的统计特性。这就导致模型没有精力去学习真正有用的特征表示训练效率非常低。如果我们用更高的学习率去加快训练分布变化会更剧烈后层就更难跟上直接导致收敛困难甚至发散。这也是为什么在 BN 出现之前人们训深层网络时学习率要设得非常保守。第二容易卡在激活函数的饱和区。拿经典的 Sigmoid 函数来说它只在输入靠近 0 的区间有较大的梯度一旦输入的绝对值偏大函数就会进入饱和区梯度趋近于 0。深层网络在更新过程中如果某层输入的数值分布逐渐偏移到比较大的区间那么这一层的神经元就很容易进入饱和状态梯度几乎传递不下去。你可以想象一下一个神经网络的大部分神经元都在“装死”反向传播的信号每过一层都衰减一点最后传到前面几层时已经所剩无几。这就解释了为什么没有 BN 的深层网络经常出现“前面几层学不动”的现象。第三对参数初始化和学习率极其敏感。在 ICS 问题严重的网络里一个稍微差一点的初始化方式或者一次稍大的参数更新都可能让某层输入分布发生剧烈变化进而引发雪崩式的梯度问题。这也是为什么早期训练深度模型非常依赖各种精细的初始化技巧就像一个没有扶手的人在走钢丝每一步都得小心翼翼。一句话总结ICS 的核心危害不是分布本身“不均匀”而是分布一直在变导致训练环境不稳定。BN 的初衷就是要抑制这种内部变化。2. Batch Normalization 的核心设计与机制2.1 归一化加两个可学习参数BN 的计算流程拆解针对 ICS2015 年 Sergey Ioffe 和 Christian Szegedy 提出了 Batch Normalization。它的思路非常直接既然网络中间层输入分布不稳定那我就在每一层激活函数之前把数据强行拉回到一个稳定的分布区间。BN 的核心操作分两步。第一步对当前 mini-batch 内的数据进行归一化。假设某一层有 d 维输入即输入是 x (x1, x2, ..., xd)对每一维特征 k在当前 batch 上计算均值和方差μ_B (1/m) * Σ x_iσ_B² (1/m) * Σ (x_i - μ_B)²然后用这两个统计量做标准化x_hat_i (x_i - μ_B) / sqrt(σ_B² ε)这里的 m 是当前 batch 的样本数量ε 是一个很小的常数防止分母为 0。经过这一步该维特征在当前 batch 上的均值约为 0方差约为 1。但这里有个问题如果只做标准化会限制网络的表达能力。比如对某个层来说它原本学到的特征分布可能是有偏的强行拉回标准正态分布可能把有价值的分布特征也抹掉了。所以 BN 又加了第二步——引入了两个可学习的参数 γ 和 β对归一化后的结果做线性变换y_i γ * x_hat_i β如果网络需要保持原来的分布γ 可以被学成 sqrt(σ² ε)β 被学成 μ这样变换就退化为恒等变换。也就是说BN 让网络自己决定“标准化到什么程度”而不是由开发者硬性规定这就是它设计巧妙的地方。在实践里你不需要手写这些统计量。PyTorch 里一行nn.BatchNorm2d(num_features)就搞定了具体计算细节由框架帮你完成。但建议每个刚学 BN 的人都手推一遍上面的公式理解这个流程后后面遇到各种奇怪问题时你才能快速定位。2.2 训练与推理两套统计量必须分清BN 有一个很关键、也很容易踩坑的细节训练阶段和推理阶段使用的统计量不是一套。训练阶段epoch 内每个 step 都会根据当前 mini-batch 的样本计算均值 μ_B 和方差 σ_B²用它们来归一化数据。同时网络还会用滑动平均的方式维护一组全局统计量running_mean 和 running_var。每次更新时running_mean (1 - momentum) * running_mean momentum * μ_B。PyTorch 中默认的 momentum 是 0.1这个值表示当前 batch 的统计量有多大的权重进入全局统计量。推理阶段我们不再计算当前 batch 的统计量而是直接使用训练时维护好的 running_mean 和 running_var 来做归一化。这样做的原因很简单推理时可能一次只来一条样本样本量为 1 时 batch 统计量没有任何统计意义。而且推理时我们希望结果是确定性的不能因为输入顺序不同导致同一条样本得到不同的输出。在代码层面PyTorch 通过model.train()和model.eval()来切换这两种状态。我见过太多新手在这上面翻车模型训练得挺好的验证时忘了加model.eval()结果推理出来结果完全不对。背后的原因就是 BN 层在两种模式下走了完全不同的分支。2.3 关于 BN 有效性的另类解释它让损失曲面变平滑了关于 BN 为什么有效有一个很流行的“标准答案”它解决了 Internal Covariate Shift。但在 2018 年Google 的研究者通过大量实验对比发现事情没这么简单。他们发现即使使用 BN 之后网络内部的输入分布仍然在变化但训练照样很稳定。真正重要的原因是BN 重塑了损失函数的曲面让优化问题变得更温和、更容易收敛。你可以这样理解在没有 BN 的网络里损失函数曲面高低起伏剧烈像一座到处是悬崖和尖峰的山脉梯度方向突变一不小心就掉进坑里。而加了 BN 之后损失曲面被大幅“磨平”了优化路径更加平缓梯度信号也更稳定。论文里用了一个非常直观的指标Lipschitz 常数。简单说就是这个常数衡量了函数变化有多剧烈BN 能显著减小这个常数让损失函数对参数更新不那么敏感。这带来了三个连锁好处一是可以放心大胆使用更大的学习率训练速度大幅提升二是模型对权重初始化的依赖变低了不用再那么小心翼翼地挑选初始化方案三是对正则化有一定帮助因为每个 mini-batch 的均值和方差略有差异相当于给训练过程引入了轻微噪声这种噪声在有些任务上起到了类似 Dropout 的正则化效果。所以在面试或和别人聊 BN 的时候不要只会背“解决 ICS”这一条。能把“平滑损失曲面”这个层面讲清楚才说明你是真的理解 BN。3. PyTorch 中落地 BN从模型搭建到参数调优3.1 手写一个简化版 BN彻底搞懂内部逻辑虽然 PyTorch 自带完善的 BN 实现但我强烈建议你手写一个简化版哪怕只是在纸上跑通逻辑也行。对于理解 BN 的内部机制这个步骤的价值远高于反复看文档。下面是我常用的一个简化版 BN 实现核心就是复现训练和推理的完整逻辑import torch import torch.nn as nn class MyBatchNorm(nn.Module): def __init__(self, num_features, eps1e-5, momentum0.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)) self.training True def forward(self, x): # 这里假设输入 x 的形状为 [N, C, H, W] N, C, H, W x.shape # 对每个通道独立计算统计量 x_reshaped x.permute(1, 0, 2, 3).reshape(C, -1) if self.training: # 在当前 batch 内计算每个通道的均值、方差 mean x_reshaped.mean(dim1) var x_reshaped.var(dim1, unbiasedFalse) # 更新滑动统计量 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 # 归一化 normed (x_reshaped - mean.view(-1, 1)) / torch.sqrt(var.view(-1, 1) self.eps) # 缩放和平移可学习 out self.gamma.view(-1, 1) * normed self.beta.view(-1, 1) out out.reshape(C, N, H, W).permute(1, 0, 2, 3) return out这段代码里有两个细节值得圈出来。第一个是var计算时用了unbiasedFalse也就是除以 n 而不是 n-1这是和 BN 原论文保持一致的做法。第二个是滑动统计量用register_buffer注册这样模型保存和加载时running_mean 和 running_var 会跟着一起走不会因为被当成普通 Python 属性而丢失。用这个手写版和nn.BatchNorm2d做对比在同样输入下输出几乎一致用默认的 affineTrue、momentum0.1 对齐参数你还可以随机初始化自建 BN 的参数一一对应。跑通一遍之后你对 BN 的“训练时用 batch 统计量、推理时用全局统计量”这套机制会产生肌肉记忆后面遇到 BN 在推理时表现异常的问题第一反应就知道该查哪个开关。3.2 在经典 CNN 上做有无 BN 的对比实验说再多理论不如实际跑一个对比实验。我直接给出一个可以快速验证的模板在 CIFAR-10 上用一个 6 层左右的 CNN分别训练带 BN 和不带 BN 的版本观察两者的训练曲线差异。import torch import torch.nn as nn import torch.optim as optim import torchvision import torchvision.transforms as transforms device torch.device(cuda if torch.cuda.is_available() else cpu) # 数据准备注意这里的归一化只是把数据搬到 0-1 附近 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) ]) trainset torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform) trainloader torch.utils.data.DataLoader(trainset, batch_size128, shuffleTrue, num_workers2) def make_cnn(use_bn): layers [] in_channels 3 for out_channels in [32, 64, 128]: layers.append(nn.Conv2d(in_channels, out_channels, 3, padding1)) if use_bn: layers.append(nn.BatchNorm2d(out_channels)) layers.append(nn.ReLU(inplaceTrue)) layers.append(nn.MaxPool2d(2)) in_channels out_channels layers.append(nn.Flatten()) layers.append(nn.Linear(in_channels * 4 * 4, 256)) if use_bn: layers.append(nn.BatchNorm1d(256)) layers.append(nn.ReLU(inplaceTrue)) layers.append(nn.Linear(256, 10)) return nn.Sequential(*layers) model_bn make_cnn(use_bnTrue).to(device) model_plain make_cnn(use_bnFalse).to(device) # 很关键的一点带 BN 的模型可以使用更大的学习率 optimizer_bn optim.SGD(model_bn.parameters(), lr0.1, momentum0.9, weight_decay5e-4) optimizer_plain optim.SGD(model_plain.parameters(), lr0.01, momentum0.9, weight_decay5e-4)实际训练时你会发现几件事。第一带 BN 的网络在前几个 epoch 的 loss 下降速度明显更快。第二不带 BN 的网络一旦把学习率调到 0.1很快就会发散而带 BN 的模型在 lr0.1 下依然稳健甚至还能往上加。第三带 BN 的模型最终收敛精度会高一些这部分来自更稳定的训练过程也有一部分来自 BN 带来的正则化效果。这个实验用 CPU 跑都能在几分钟内看到趋势非常适合用来向身边人展示 BN 的价值。3.3 BN 超参与使用技巧BN 层看起来就两个可学习参数加一套滑动统计量但实际调参时还是有不少门道。我整理了一张常用参数与我的经验设置方便你抄作业。参数默认值经验说明eps1e-5保持默认即可。数值稳定用一般不需要动。momentum0.1训练数据多、batch 多的时候可以适当调小到 0.01让滑动统计量更平滑。affineTrue一般保持 True。如果为了省显存可以关掉但会牺牲表达能力。track_running_statsTrue保持 True。如果设成 False推理时会用当前 batch 统计量结果不稳定。除了这些参数还有几个使用位置上的经验。BN 一般是放在卷积层或全连接层之后、激活函数之前。顺序是Conv - BN - ReLU。不要先把 ReLU 放在 BN 前面因为 ReLU 会截断负值导致归一化时的数据分布和原始输出不一致BN 的效果会受影响。大批量训练时BN 的 batch 统计量更精准效果更好小批量时统计量噪声大训练可能更不稳定这点在下一节会详细展开。在残差网络里BN 通常放在卷积之后、残差相加之前。分支内部的 BN 都是为了稳定分支内部的信号流动而相加之后一般不会立刻再放一个 BN否则会破坏恒等映射的快捷路径。4. 实际训练中那些与 BN 相关的经典翻车现场4.1 常见问题速查表训练跑多了你会发现 BN 相关的问题基本都集中在几个固定场景里。我在下面列成了表格方便直接对照排查。现象可能原因解决办法训练正常推理结果完全崩掉忘了切换 model.eval()BN 用了有噪声的 batch 统计量推理前调用 model.eval()另外断点续训时也要恢复正确状态batch size 很小loss 震荡严重mini-batch 统计量噪声太大BN 不稳定用较大的 batch或换成 GroupNorm 等与 batch 无关的归一化方法训练后期验证集 loss 反而升高BN 的滑动统计量更新太慢跟不上模型分布调小 momentum降低后期学习率给统计量更多时间追上模型刚开始训练就出现 NaNBN 在归一化时方差为 0或梯度爆炸检查 eps 是否过小尽量使用默认 eps观察梯度范数迁移学习时新任务效果差预训练模型的 running_mean/running_var 与新领域数据分布不匹配微调初期可以冻结 BN 统计量只训练其他参数或在新数据上做少量前向更新统计量这里面最隐蔽的一个问题就是“预训练模型 fine-tune 时 BN 参数怎么处理”。很多做迁移学习的同学直接把带 BN 的预训练模型在新数据集上从头训练结果发现效果甚至不如随机初始化的小模型。原因就在于新数据的分布和原预训练数据差异较大模型一开始 forward 时running_mean 和 running_var 还是旧的相当于每层的输入被一个错误分布做了归一化梯度方向乱七八糟。我常用的做法是微调初期把 BN 参数设为冻结状态只训练迁移任务新增的层等模型在新数据上跑出一定准确率后再解冻 BN 层做联合微调。解冻后还需要把学习率调小一点防止 BN 统计量在微调初期被大梯度破坏掉。4.2 小 batch size 场景下的替代方案BN 的一个致命前提是 batch 足够大统计量才有意义。但在目标检测、点云分割这类任务里单卡 batch size 经常只有 2 或 4这时 BN 几乎是在“裸奔”。如果你在一个 batch 里看到某个通道的方差接近 0基本就是统计量失真了归一化出来的数值会非常极端训练直接爆掉。针对小 batch 场景有几个替换方案可以考虑。第一个是 SyncBN也就是跨卡同步的 Batch Normalization。它把多张 GPU 上的同一 batch 级联起来计算统计量单卡 batch 小没关系只要总 batch 够大就行。PyTorch 里用torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)可以把普通 BN 一键转换成 SyncBN。注意它依赖分布式训练环境单卡模式下没有意义。第二个是 GroupNorm它对每个样本独立地在通道维度上分组做归一化完全不依赖 batch 大小。这在大 batch 和小 batch 下表现都比较稳定常用于检测、分割等显存紧张的任务。很多经典检测框架把 Backbone 的 BN 替换成 GroupNorm 后都取得了不错的效果。第三个是 LayerNorm 和 InstanceNorm一个按样本和全部通道做归一化一个按样本和单通道做归一化。它们在自然语言处理、图像风格迁移等任务里更常用。总体原则是如果你做视觉任务且 batch 稳定在 32 以上继续用 BN如果 batch 很小又不想折腾 SyncBNGroupNorm 是性价比最高的选择。4.3 BN 与 Dropout、归一化变体的搭配经验关于 BN 和 Dropout 能不能一起用这个争论从 BN 一出来就存在。我的个人经验是在卷积网络里BN 带来的正则化已经比较强再叠加 Dropout 往往没有必要甚至可能因为双重正则化导致欠拟合。如果你使用 ResNet 这类带 BN 的骨干网络最后一层全连接之前不加 Dropout 也很稳。但在全连接层较多的网络里如果你发现验证集 loss 明显高于训练集 loss适当在 BN 之后、下一个线性层之前加一个 Dropout还是能带来一些提升的。还有一个需要注意的点是 BN 和 L2 正则化的配合。BN 把数据归一化之后权重衰减的作用对象、幅度都发生了变化。实测下来使用 BN 的模型通常会把 weight_decay 调大一些比如从 5e-4 调到 1e-3反而能获得更好的泛化表现。这个规律在 ResNet 系列上表现得很明显只训不带 BN 的小模型时weight_decay 的影响就没那么显著。另外有过训练 NLP 模型经验的人肯定会觉得 BN 在 Transformer 里并不好用这是正常的。BN 在图像任务里好使很大程度上是因为 CNN 的卷积特征在通道维度上存在稳定的统计规律。而在 NLP 中不同样本的序列长度差异大一个 batch 内的统计量受长度影响严重BN 的效果往往不如 LayerNorm。所以不要盲目把 BN 搬到所有领域理解归一化的本质是调整分布再根据数据特性选择具体方式才是正道。写在最后我个人在实际训练里对 BN 最深的感受就是它像给网络加了一根“安全绳”。以前训深层网络每一步都要小心翼翼初始化差一点、学习率大一点都可能翻车有了 BN 之后很多巧合和不稳定因素被直接抹平了。它没有让你的模型突然变得更强但它让“训练一个深层模型”这件事本身变得可靠了。当然BN 远不是万能的。它依赖 batch 大小、需要维护滑动统计量、在不同领域里的适配性也不一样。但作为深度学习最基础、最核心的归一化方法之一把它的原理和实操细节吃透你会发现自己调试网络的能力会上一个台阶。尤其是当你遇到模型训练不稳定的问题时只要从数据分布这个角度去排查很多看似玄学的问题都会有迹可循。
返回列表