\[ \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} \]
批量归一化:参考答案
返回正文练习 · 返回答案索引
说明
以下答案与正文 9 道题逐题对应。推导题给出主要步骤;编程和实验题给出参考程序和检查方法,并说明结果适用于哪些条件。
第一式的分子求和为 \(\sum_i(z_i-\mu_B)=\sum_i z_i-m\mu_B=0\),分母对同一特征是常数,因此标准化均值为 0。第二式为 \(m^{-1}\sum_i(z_i-\mu_B)^2/(v_B+\epsilon)=v_B/(v_B+\epsilon)\)。当 \(\epsilon>0\) 时平方均值略小于 1;只有忽略 \(\epsilon\) 且 \(v_B>0\) 时才等于 1。
批均值为 \(\mu_B=2\),批方差为 \([(1-2)^2+(3-2)^2]/2=1\)。标准化结果为 \((-1,1)\),经过尺度与平移后输出为 \((2(-1)+1,2(1)+1)=(-1,3)\)。实际程序必须使用正的 \(\epsilon\);题中取 0 只是为了得到便于手算的精确值。
令 \(\widehat g_i=g_i\gamma\)。则 \(\partial\mathcal J/\partial\beta=\sum_i g_i\)、\(\partial\mathcal J/\partial\gamma=\sum_i g_i\widehat z_i\),并有 \(\partial\mathcal J/\partial z_i=[m\widehat g_i-\sum_j\widehat g_j-\widehat z_i\sum_j(\widehat g_j\widehat z_j)]/[m\sqrt{v_B+\epsilon}]\)。对多特征批次,所有求和沿批次轴逐特征进行;输入梯度与 \(\bZ\) 同形,\(\bgamma\) 和 \(\bbeta\) 的梯度与各自参数同形。
第一次更新为 \(0.9(0)+0.1(2)=0.2\),第二次为 \(0.9(0.2)+0.1(4)=0.58\)。推断时使用冻结运行量,可使某个样本的预测不依赖恰好与它组成批次的其他样本,并允许稳定的单样本预测。运行统计量只由训练批次更新,验证和测试阶段不得继续改变。
二维输入可按特征轴保存运行量:
def bn_forward(z, gamma, beta, running, training, rho=0.9, eps=1e-5):
if z.ndim != 2 or z.shape[1] != gamma.size:
raise ValueError("输入或参数维度错误")
if training:
mean = z.mean(axis=0)
var = z.var(axis=0, ddof=0)
zhat = (z - mean) / np.sqrt(var + eps)
running["mean"] = rho * running["mean"] + (1-rho) * mean
running["var"] = rho * running["var"] + (1-rho) * var
cache = (zhat, gamma, var, eps)
else:
zhat = (z - running["mean"]) / np.sqrt(running["var"] + eps)
cache = None
return gamma * zhat + beta, cache
运行均值和方差的初始化方式以及方差估计约定必须与所对照的框架一致。
中心差分 对每个元素使用 \([\mathcal J(x+h)-\mathcal J(x-h)]/(2h)\);选择若干 \(h\) 检查误差先下降后受舍入影响。与 PyTorch 自动微分 对照时,必须使用相同的批方差约定、\(\epsilon\)、\(\gamma\) 和 \(\beta\)。float64 小张量上,手工梯度与自动微分应在严格容差内一致;若只有输入梯度错误,应先检查批次轴求和与缓存是否来自同一次前向。
训练模式中使用 out_before_scale 或缓存的 zhat 检查逐列均值及平方均值;常数列在正 \(\epsilon\) 下应输出有限值;推断时把目标样本分别与两个不同批次拼接,预测必须相同;dz.shape == z.shape 且参数梯度维度等于特征数;输入列数与运行量不一致时抛出异常。还应验证推断调用不会修改运行统计量。
为了使比较结果可靠,各组只改变是否使用 BN,其他条件保持相同;若 BN 增加 \(\gamma,\beta\),应如实计入参数量。BN 可能使训练过程更稳定或更快,但不能预先认定它一定会提高最终指标。应针对每个随机种子报告测试指标及其波动、达到同一验证指标所需时间、总训练时间、程序预先运行后测得的固定批量预测时间和峰值内存。预测时 BN 使用训练阶段得到的统计量,会增加少量逐元素运算;结论应以多次计时结果为准。
批量大小改变了批统计噪声,也改变硬件利用率;梯度累积只能对齐每次更新的样本数,不能让 BN 看到同样的统计批次。每组固定更新次数、种子、划分和学习率规则,报告运行均值与全训练集均值的距离、任务指标、训练秒数、固定批量预测时间及内存。模型参数量不随批量大小改变;峰值内存通常随批量增大,训练时间可能随批量大小发生变化。