使用BatchNorm偏置填充边界:确保推理一致性与数值稳定性
1. 问题背景在模型推理、结构重参数化或算子融合时可能需要调整 BatchNorm 和 Padding 的执行顺序。原来的计算顺序可能是xzero_pad(x)xbn(x)调整后变成xbn(x)xpad(x)这时第二种写法中的 Padding 通常不能继续填充 0而应当填充BN(0)\mathrm{BN}(0)BN(0)即输入值 0 经过 BatchNorm 后得到的结果。本文推导这个填充值并说明它真正解决的是什么问题。2. 推理阶段的 BatchNorm对于输入xxxBatchNorm 的计算公式为yγx−μσ2εβy\gamma\frac{x-\mu}{\sqrt{\sigma^2\varepsilon}}\betayγσ2εx−μβ。其中μ\muμ是运行均值σ2\sigma^2σ2是运行方差γ\gammaγ是缩放参数β\betaβ是平移参数ε\varepsilonε是保证数值稳定性的常数。在推理阶段μ\muμ和σ2\sigma^2σ2都是固定值因此可以展开为yγσ2εxβ−γμσ2εy\frac{\gamma}{\sqrt{\sigma^2\varepsilon}}x\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}yσ2εγxβ−σ2εγμ。令aγσ2εa\frac{\gamma}{\sqrt{\sigma^2\varepsilon}}aσ2εγbβ−γμσ2εb\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}bβ−σ2εγμ那么 BatchNorm 可以写成yaxbyaxbyaxb也就是说推理阶段的 BatchNorm 本质上是一个固定的逐通道仿射变换。3. BatchNorm 中的 BN(0)令输入x0x0x0代入 BatchNorm 公式BN(0)γ0−μσ2εβ\mathrm{BN}(0)\gamma\frac{0-\mu}{\sqrt{\sigma^2\varepsilon}}\betaBN(0)γσ2ε0−μβ整理后得到BN(0)β−γμσ2ε\mathrm{BN}(0)\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}BN(0)β−σ2εγμ因此BN(0)\mathrm{BN}(0)BN(0)就是 BatchNorm 的等效偏置bbb。需要注意它不一定等于bn.bias。bn.bias只对应公式中的β\betaβ而完整的等效偏置还包含均值、方差和缩放参数bβ−γμσ2εb\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}bβ−σ2εγμ对应的 PyTorch 实现如下importtorchdefget_bn_zero_value(bn):return(bn.bias-bn.running_mean*bn.weight/torch.sqrt(bn.running_varbn.eps))4. 为什么 BatchNorm 后不能直接补 0假设Pad0(x)\mathrm{Pad}_0(x)Pad0(x)表示使用 0 对输入进行填充。先补 0再执行 BatchNorm计算过程为BN(Pad0(x))\mathrm{BN}(\mathrm{Pad}_0(x))BN(Pad0(x))Padding 增加的边界值虽然最开始是 0但这些 0 还会经过 BatchNorm因此最终边界值会变成BN(0)\mathrm{BN}(0)BN(0)。先执行 BatchNorm再补 0计算过程为Pad0(BN(x))\mathrm{Pad}_0(\mathrm{BN}(x))Pad0(BN(x))此时 Padding 位于 BatchNorm 后面新增加的边界值仍然是 0。因此一般情况下BN(Pad0(x))≠Pad0(BN(x))\mathrm{BN}(\mathrm{Pad}_0(x))\ne\mathrm{Pad}_0(\mathrm{BN}(x))BN(Pad0(x))Pad0(BN(x))要想让两种执行顺序保持一致BatchNorm 后的 Padding 值必须设置为BN(0)\mathrm{BN}(0)BN(0)BN(Pad0(x))PadBN(0)(BN(x))\mathrm{BN}(\mathrm{Pad}_0(x))\mathrm{Pad}_{\mathrm{BN}(0)}(\mathrm{BN}(x))BN(Pad0(x))PadBN(0)(BN(x))其中BN(0)β−γμσ2ε\mathrm{BN}(0)\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}BN(0)β−σ2εγμ这就是正确边界填充值的来源。5. 一个简单例子假设推理阶段的 BatchNorm 可以表示为BN(x)2x3\mathrm{BN}(x)2x3BN(x)2x3那么BN(0)3\mathrm{BN}(0)3BN(0)3假设输入为[1,2,3][1,2,3][1,2,3]。先补 0再执行 BatchNorm[0,1,2,3,0]→[3,5,7,9,3][0,1,2,3,0]\rightarrow[3,5,7,9,3][0,1,2,3,0]→[3,5,7,9,3]如果先执行 BatchNorm再补 0[1,2,3]→[5,7,9]→[0,5,7,9,0][1,2,3]\rightarrow[5,7,9]\rightarrow[0,5,7,9,0][1,2,3]→[5,7,9]→[0,5,7,9,0]两者显然不相等。如果在 BatchNorm 后填充BN(0)3\mathrm{BN}(0)3BN(0)3[5,7,9]→[3,5,7,9,3][5,7,9]\rightarrow[3,5,7,9,3][5,7,9]→[3,5,7,9,3]此时两种计算顺序的结果完全一致。6. BatchNorm2d 需要按通道填充对于BatchNorm2d每个通道都有独立的γ\gammaγ、β\betaβ、μ\muμ和σ2\sigma^2σ2。因此第ccc个通道对应的填充值为BNc(0)βc−γcμcσc2ε\mathrm{BN}_c(0)\beta_c-\frac{\gamma_c\mu_c}{\sqrt{\sigma_c^2\varepsilon}}BNc(0)βc−σc2εγcμc不同通道的填充值通常不同不能简单使用一个标量填充所有通道。假设特征图形状为[N, C, H, W]计算得到的BN(0)形状为[C]使用时通常需要整理为[1, C, 1, 1]以便按通道广播。一种简单实现如下defpad_by_channel(x,value,padding1):n,c,h,wx.shape valuevalue.to(devicex.device,dtypex.dtype).view(1,c,1,1)outvalue.expand(n,c,h2*padding,w2*padding).clone()out[:,:,padding:paddingh,padding:paddingw]xreturnout使用方式bn.eval()withtorch.no_grad():ybn(x)padding_valueget_bn_zero_value(bn)ypad_by_channel(y,padding_value)7. 使用条件上述等价关系主要适用于推理阶段使用前应执行model.eval()在eval模式下BatchNorm 使用固定的running_mean和running_var因此可以写成固定的仿射变换yaxbyaxbyaxb。训练模式下BatchNorm 通常使用当前 batch 的均值和方差。先 Padding 会改变参与统计的数据因此一般不能直接使用上述等价变换。如果 BatchNorm 设置了affineFalse则可以认为γ1\gamma1γ1、β0\beta0β0。如果设置了track_running_statsFalseBatchNorm 在推理时仍可能依赖当前输入的统计量此时也不适合使用固定的BN(0)\mathrm{BN}(0)BN(0)。8. 这不是在修复边界统计问题使用BN(0)\mathrm{BN}(0)BN(0)填充并不是因为“边界缺少邻域数据导致 BatchNorm 计算不准确”。BatchNorm 不是局部窗口运算。对于BatchNorm2d训练时通常按通道在N×H×WN\times H\times WN×H×W范围内计算统计量边缘位置和中间位置使用同一组通道统计参数。边界缺少邻域是卷积运算需要考虑的问题不是 BatchNorm 本身的统计方式。因此填充BN(0)\mathrm{BN}(0)BN(0)的真正目的只有一个在调整 BatchNorm 和 Padding 的执行顺序后保持变换前后的计算结果一致。核心等价关系为BN(Pad0(x))PadBN(0)(BN(x))\mathrm{BN}(\mathrm{Pad}_0(x))\mathrm{Pad}_{\mathrm{BN}(0)}(\mathrm{BN}(x))BN(Pad0(x))PadBN(0)(BN(x))核心填充值为BN(0)β−γμσ2ε\mathrm{BN}(0)\beta-\frac{\gamma\mu}{\sqrt{\sigma^2\varepsilon}}BN(0)β−σ2εγμ