\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

批量归一化(Batch Normalization)#

学习目标与记号#

  1. 准确解释批次统计量、标准化以及可学习的尺度和平移;

  2. 根据公式和张量维度分析运行统计量,并识别常见实现错误;

  3. Python 实现或验证批量归一化的后向传播;在其他条件相同的情况下,只改变一个需要研究的因素进行比较,并根据实验结果分析该因素可能带来的影响。

  本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;层编号写作上标 \([l]\)⁠,样本或时间编号写作下标;转置写作 \(\trans\)⁠。除非另有说明,批量样本按行存放。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 批量归一化答案⁠。

  本节默认读者熟悉 多样本向量化后向传播链式法则⁠;批量归一化与 Dropout 的作用机制不同,不应混为同一种正则化方法。

  Ioffe 与 Szegedy(2015) 提出的批量归一化(Batch Normalization,BN)在一个小批量内按特征维度标准化中间变量,再用可学习参数恢复合适的尺度和位置。它通常能允许较大的学习率、降低训练对初始化的敏感程度,并改善深层网络的优化。该论文将其作用解释为缓解“内部协变量偏移”,但 Santurkar 等(2018) 等后续研究表明,这并不是理解其效果的唯一机制;更稳妥的表述是,BN 通过对每个小批量的中间结果进行标准化,改变了参数更新对模型输出和损失函数的影响,使不同层的数值尺度和梯度变化通常更加稳定,因此模型往往更容易训练。

训练阶段的前向传播#

  采用“样本按行堆叠”的约定。设第 \(l\) 层当前小批量大小为 \(m\)⁠,仿射变换输出为

\[\bZ^{[l]} =[(\bz_1^{[l]})\trans;\ldots;(\bz_m^{[l]})\trans] =\bA^{[l-1]}\cdot\left(\bW^{[l]}\right)\trans +\bone\cdot\left(\bb^{[l]}\right)\trans \in\mathbb{R}^{m\times d^{[l]}},\]

其中 \(\bone\in\mathbb{R}^{m\times1}\)⁠。BN 对 \(\bZ^{[l]}\) 的每一列分别计算小批量均值与方差:

(80)#\[\begin{split}\begin{aligned} \bmu_B^{[l]} &=\frac{1}{m}\left(\bZ^{[l]}\right)\trans\cdot\bone, \\ \bv_B^{[l]} &=\frac{1}{m}\sum_{i=1}^{m} \left(\bz_i^{[l]}-\bmu_B^{[l]}\right) \odot \left(\bz_i^{[l]}-\bmu_B^{[l]}\right), \\ \widehat{\bz}_i^{[l]} &=\frac{\bz_i^{[l]}-\bmu_B^{[l]}} {\sqrt{\bv_B^{[l]}+\epsilon}}, \\ \widetilde{\bz}_i^{[l]} &=\bgamma^{[l]}\odot\widehat{\bz}_i^{[l]} +\bbeta^{[l]}, \\ \ba_i^{[l]} &=g^{[l]}\!\left(\widetilde{\bz}_i^{[l]}\right), \end{aligned}\end{split}\]

其中平方根、除法和 \(\odot\) 均逐元素计算,\(\epsilon>0\) 用于避免除零。\(\bgamma^{[l]}\)\(\bbeta^{[l]}\) 是通过后向传播学习的尺度和平移参数。标准化后的变量在当前小批量内均值接近 0、方差接近 1;变换后的均值为 \(\bbeta^{[l]}\)⁠,方差近似为 \((\bgamma^{[l]})^2\)⁠,而不应被描述为“标准差等于 \(\bgamma^{[l]}\)⁠”——后者在 \(\gamma_j<0\) 时尤其不成立。

  将各样本按行堆叠为 \(\widehat{\bZ}^{[l]}=[(\widehat{\bz}_1^{[l]})\trans;\ldots;(\widehat{\bz}_m^{[l]})\trans]\)⁠、\(\widetilde{\bZ}^{[l]}=[(\widetilde{\bz}_1^{[l]})\trans;\ldots;(\widetilde{\bz}_m^{[l]})\trans]\)\(\bA^{[l]}=[(\ba_1^{[l]})\trans;\ldots;(\ba_m^{[l]})\trans]\) 后,可用广播写成紧凑的矩阵形式:

(81)#\[\begin{split}\begin{aligned} \bZ_c^{[l]} &=\bZ^{[l]}-\bone\cdot\left(\bmu_B^{[l]}\right)\trans,\\ \widehat{\bZ}^{[l]} &=\bZ_c^{[l]}\odot \left[\bone\cdot\left\{\left(\bv_B^{[l]}+\epsilon\right)^{-1/2}\right\}\trans\right],\\ \widetilde{\bZ}^{[l]} &=\widehat{\bZ}^{[l]}\odot \left\{\bone\cdot(\bgamma^{[l]})\trans\right\} +\bone\cdot(\bbeta^{[l]})\trans. \end{aligned}\end{split}\]

  在“仿射层—BN”组合中,批量中心化会抵消仿射层偏置 \(\bb^{[l]}\) 的作用,所以实现中常令该仿射层不使用偏置,并由 \(\bbeta^{[l]}\) 承担平移。BN 放在激活函数之前还是之后属于架构选择;常见做法是“线性/卷积—BN—激活”,但应遵循所用架构的既定设计。

BN 的轻微正则化作用

  训练阶段使用当前小批量的均值和方差。同一个样本若与不同样本组成小批量,得到的均值和方差通常略有不同,因而其标准化结果和后续梯度也会出现轻微波动。这种随机波动相当于在训练过程中加入少量噪声,可能减轻模型对训练数据细节的过度拟合,因此 BN 有时会表现出轻微的正则化作用。

  这种作用的强弱与批量大小、批量中的样本组成以及具体任务有关。批量较小时统计量的波动通常更明显,但批量过小也可能使训练不稳定,因此不能简单地通过减小批量来增强正则化。推断阶段使用训练过程中得到并固定下来的运行统计量,上述随机波动不再存在。BN 的主要目的仍是改善训练过程,不能保证替代权重衰减、数据增强或 Dropout⁠。

训练阶段与推断阶段#

  训练阶段使用当前小批量统计量,因此一个样本的输出会受到同批其他样本影响。与此同时,BN 层维护总体均值和方差的滑动估计。例如可写为

\[\begin{split}\begin{aligned} \bmu_{\mathrm{run}} &\leftarrow \rho\bmu_{\mathrm{run}}+(1-\rho)\bmu_B,\\ \bv_{\mathrm{run}} &\leftarrow \rho\bv_{\mathrm{run}}+(1-\rho)\bv_B, \end{aligned}\end{split}\]

其中 \(\rho\) 是动量系数。不同框架对 momentum 参数的命名方向和方差校正细节可能不同,应以所用 API 为准。

  推断阶段不能依赖“当前批量”,否则同一个样本会因与谁一起预测而得到不同结果。此时关闭批次统计,使用训练阶段积累的 \(\bmu_{\mathrm{run}}\)\(\bv_{\mathrm{run}}\)⁠:

\[\widehat{\bz} =\frac{\bz-\bmu_{\mathrm{run}}} {\sqrt{\bv_{\mathrm{run}}+\epsilon}}, \qquad \widetilde{\bz}=\bgamma\odot\widehat{\bz}+\bbeta.\]

  这就是深度学习框架中 train()eval() 状态对 BN 至关重要的原因。验证和测试时若忘记切换到评估模式,指标可能不稳定甚至产生系统偏差。

后向传播#

  设上游梯度 \(\bG=\partial\mathcal{J}/\partial\widetilde{\bZ}\in\mathbb{R}^{m\times d}\)⁠。尺度和平移参数的梯度为

\[\frac{\partial\mathcal{J}}{\partial\bgamma} =(\bG\odot\widehat{\bZ})\trans\cdot\bone, \qquad \frac{\partial\mathcal{J}}{\partial\bbeta} =\bG\trans\cdot\bone.\]

  令 \(\bH=\bG\odot\{\bone\cdot\bgamma\trans\}\)⁠、\(\bs=(\bv_B+\epsilon)^{-1/2}\)⁠,其中幂运算逐元素执行。对标准化输入的完整链式求导可以化简为

\[\frac{\partial\mathcal{J}}{\partial\bZ} =\frac{1}{m} \left(\bone\cdot\bs\trans\right) \odot \left[ m\bH -\bone\cdot(\bone\trans\cdot\bH) -\widehat{\bZ}\odot \left\{\bone\cdot\left(\bone\trans\cdot (\bH\odot\widehat{\bZ})\right)\right\} \right].\]

  为避免转置造成歧义,也可以逐列理解该式:每一列的梯度都要减去该列梯度的均值,以及它在标准化变量方向上的分量。得到 \(\partial\mathcal{J}/\partial\bZ\) 后,仿射层梯度仍为

\[\frac{\partial\mathcal{J}}{\partial\bW} =\left(\frac{\partial\mathcal{J}}{\partial\bZ}\right)\trans\cdot\bA^{[l-1]}, \qquad \frac{\partial\mathcal{J}}{\partial\bA^{[l-1]}} =\frac{\partial\mathcal{J}}{\partial\bZ}\cdot\bW.\]

  实际编程时应使用自动微分;手工公式主要用于理解每个分支为何都必须参与链式法则。

局限与使用建议#

  1. 很小的批量会使均值和方差估计噪声较大。可考虑冻结统计量、累积统计量,或改用 Layer Normalization、Group Normalization 等不依赖批量维度的方法。我们将在后续内容中对 Layer Normalization 进行详细介绍。

  2. 分布式训练需要明确统计量是在单个设备还是多个设备间同步计算;二者可能得到不同结果。

  3. BN 产生的随机性有时带来轻微正则化,但不能保证替代权重衰减、数据增强或 Dropout。

  4. 微调预训练模型时,应根据数据量和分布变化决定冻结 BN、只更新仿射参数,还是重新估计运行统计量。

Shiny 交互演示:批量归一化

  下面的交互页面用于比较使用和不使用批量归一化时的训练过程,并观察批量统计量与运行统计量。训练阶段按照小批量计算均值和方差;验证、测试以及新样本预测均切换到评估模式,使用训练阶段记录的运行统计量。

点击打开“批量归一化”交互演示

核心推导与实现核验#

核心关系

\[\widehat{\bz}_i=(\bz_i-\bmu_B)\oslash\sqrt{\bv_B+\epsilon},\]

其中 \(\oslash\) 表示逐元素相除。

  推导路径。 先对每个特征在批量轴计算均值和方差,中心化后按标准差逐元素缩放,再用可学习 \(\bgamma,\bbeta\) 恢复表示能力。

关键条件

  BN 会使批内标准化结果的均值严格等于 0。由于计算时在方差中加入了防止除以 0 的小正数 \(\epsilon\)⁠,其方差为 \(v/(v+\epsilon)\)⁠,通常接近 1,但不一定严格等于 1;当批内方差 \(v\) 很小时,两者的差别可能较为明显。

数据规模

  一次输入 \(m\) 个样本且每个样本有 \(d\) 个特征时,输入数据的大小为 \(m\times d\)⁠。均值 \(\bmu_B\)⁠、方差 \(\bv_B\)⁠、缩放参数 \(\bgamma\) 和平移参数 \(\bbeta\) 都各有 \(d\) 个数,并对同一列的 \(m\) 个样本共同使用。

常见误区

  小批量过小时统计量噪声大;忘记切换 eval 模式会使单样本预测依赖同批其他样本。

动手检查

  训练模式检查标准化值的批内均值接近 0、平方均值接近 v/(v+eps)并把手写后向计算与 自动微分 逐项比较。

数值稳定性与规模

  计算批量均值和方差时,应优先使用框架提供的可靠函数,或者先计算均值,再计算各数值与均值之差的平方平均。在方差中加入小正数 \(\epsilon\)⁠,可以避免方差为 0 或非常小时出现除以 0 或结果过大的问题。使用混合精度训练时,即使模型的部分计算采用 float16bfloat16均值和方差通常仍使用 float32 计算。模型预测新数据时,应使用训练期间逐步积累并固定下来的运行均值和运行方差,不再根据当前预测批次重新计算。

本节小结#

  1. BN 的统计轴是批量轴,参数轴是特征轴。

  2. 训练统计量与推断运行量承担不同职责。

  3. epsilon、小批量和框架方差约定会影响精确结果。

综合练习#

  程序题应固定随机种子、写出维度断言并报告运行环境;比较题还应固定数据划分、随机种子集合和训练预算。全部参考答案见 批量归一化答案⁠。

  1. 标准化性质。 对一个批次中的单个特征,设 \(\mu_B=m^{-1}\sum_{i=1}^m z_i\)⁠、\(v_B=m^{-1}\sum_{i=1}^m(z_i-\mu_B)^2\)⁠、\(\widehat z_i=(z_i-\mu_B)/\sqrt{v_B+\epsilon}\)⁠。证明 \(m^{-1}\sum_i\widehat z_i=0\)⁠,并推导 \(m^{-1}\sum_i\widehat z_i^2=v_B/(v_B+\epsilon)\)⁠。

  2. 前向具体计算。 某个特征在一个批次的两个样本中分别取值为 1 和 3,取 \(\epsilon=0\)⁠、\(\gamma=2\)⁠、\(\beta=1\)⁠。计算批均值、批方差、标准化结果以及 \(y_i=\gamma\widehat z_i+\beta\) 的两个输出。

  3. 后向传播推导。 对单个特征记上游梯度为 \(g_i=\partial\mathcal J/\partial y_i\)⁠。从计算图推导 \(\partial\mathcal J/\partial\beta\)⁠、\(\partial\mathcal J/\partial\gamma\)\(\partial\mathcal J/\partial z_i\) 的批量表达式,并核对每个梯度的维度。

  4. 运行统计量计算。 采用更新 \(\mu_{\mathrm{run}}\leftarrow\rho\mu_{\mathrm{run}}+(1-\rho)\mu_B\)⁠。若 \(\rho=0.9\)⁠、旧运行均值为 0,连续两个批均值为 2 和 4,计算两次更新后的运行均值;说明推断时为什么使用冻结的运行统计量。

  5. BN 前向实现。 用 NumPy 实现支持二维批量输入的 BN 训练与推断的前向传播的计算函数;训练时返回后向传播所需缓存并更新运行均值和运行方差,推断时只使用运行统计量。

  6. 后向与自动微分对照。 实现第 3 题的 BN 后向传播的计算,在 float64 小张量上与 PyTorch 自动微分 比较 \(\partial\mathcal J/\partial\bZ\)⁠、\(\partial\mathcal J/\partial\bgamma\)⁠、\(\partial\mathcal J/\partial\bbeta\)⁠。

  7. 测试设计。 测试第 5--6 题程序:训练输出逐特征均值接近 0、平方均值符合 \(v_B/(v_B+\epsilon)\)⁠、常数列仍产生有限值、推断结果不随同批其他样本改变、梯度与输入或参数同形,以及错误特征数被拒绝。

  8. 有无 BN 的比较。 在同一深层网络上比较不使用 BN 与使用 BN。固定训练/验证/测试划分、随机种子集合、模型宽度、优化器候选和训练预算;报告任务性能、收敛曲线、训练时间、固定批量预测时间、参数量及峰值内存。

  9. 批量大小比较。 对带 BN 的同一网络比较多个批量大小,并通过梯度累积尽量保持每次参数更新看到的样本数一致。固定数据划分、随机种子集合和参数更新次数;报告任务性能、运行统计量误差、训练时间、固定批量预测时间、参数量及峰值内存。