Pytorch中的BatchNorm2d
原理
【BatchNorm2d】是Pytorch中的归一化处理组件,其数学原理为
y = x − μ σ + ϵ γ + β \mathbf y=\frac{\mathbf x-\mu}{\sigma+\epsilon}\gamma+\beta y=σ+ϵx−μγ+β
式中, μ , σ \mu,\sigma μ,σ为期望核方差, μ = 1 m ∑ i = 1 m x i \mu=\frac{1}{m}\sum^m_{i=1}x_i μ=m1∑i=1mxi, σ 2 = 1 m ∑ i = 1 m ( x i − μ ) 2 \sigma^2=\frac{1}{m}\sum^m_{i=1}(x_i-\mu)^2 σ2=m1∑i=1m(xi−μ)2,其中 m m m是当前batch的样本数乘以特征图的空间尺寸, x i x_i xi是该通道内的所有元素。 ϵ \epsilon ϵ是为了保证数值稳定性添加的微小量, γ , β \gamma,\beta γ,β为科学系参数向量,默认 γ = 1 , β = 0 \gamma=1,\beta=0 γ=1,β=0。
在模型实现后,其输入张量的标准形状为 ( N , C , H , W ) (N,C,H,W) (N,C,H,W),分别代表批处理尺寸,通道数,高和宽。从公式可知,其输出 y \mathbf y y和输入 x \mathbf x x有着相同的形状。由于归一化针对的是通道 C C C这个维度,所以在求期望与方差的时候, m = N × H × W m=N\times H\times W m=N×H×W。
一般来说,BatchNorm2d在训练模式和推理模式下应该采取不同的行为。在训练模式下, μ , σ \mu, \sigma μ,σ为当前batch的实时统计量,同时,利用指数移动平均更新全局的运行均值(Running Mean, μ r \mu_r μr)和运行方差(Running Var, σ r \sigma_r σr),为推理阶段做准备,更新公式为
σ r = ( 1 − α ) σ r + α σ \sigma_r=(1-\alpha)\sigma_r+\alpha\sigma σr=(1−α)σr+ασ
而在推模式下,则直接使用训练阶段累积得到的全局 μ r \mu_r μr和 σ r \sigma_r σr代入标准化公式。
实现
【BatchNorm2d】的函数签名为
BatchNorm2d(num_features, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True, device=None, dtype=None)
其中
- num_features为预期输入大小 C C C
- eps 即 ϵ \epsilon ϵ,默认 10 − 5 10^{-5} 10−5
- momentum 用于 running_mean 和 running_var 计算。可以设置为 None 以进行累积移动平均(即简单平均)。默认为:0.1
- affine 默认为True,表示该模块具有可学习的仿射参数。
- track_running_stats 当设置为 True 时,此模块会跟踪运行中的均值和方差;当设置为 False 时,此模块不跟踪这些统计量,并将统计量缓冲区 running_mean 和 running_var 初始化为 None。当这些缓冲区为 None 时,此模块在训练和评估模式下始终使用批次统计量。默认值: True
测试
import torch
import torch.nn as nn
m=nn.BatchNorm2d(2,affine=True)
input=torch.randn(1,2,3,4)
output=m(input)
print(input)
print(m.weight)
print(m.bias)
print(output)
print(output.size())
其输入为
[ [ − 0.1083 0.6814 1.1752 − 0.5405 − 1.4278 − 1.8828 0.5529 2.1013 0.4706 0.8746 − 0.5013 − 0.6921 ] [ − 0.9154 1.4002 0.8023 − 0.9378 0.4337 0.2682 0.3224 0.4463 − 0.4859 1.0584 − 0.0986 − 0.6655 ] ] \begin{bmatrix} \begin{bmatrix} -0.1083& 0.6814& 1.1752&-0.5405\\ -1.4278&-1.8828& 0.5529& 2.1013\\ 0.4706& 0.8746&-0.5013&-0.6921 \end{bmatrix}& \begin{bmatrix} -0.9154& 1.4002& 0.8023&-0.9378\\ 0.4337& 0.2682& 0.3224& 0.4463\\ -0.4859& 1.0584&-0.0986&-0.6655 \end{bmatrix} \end{bmatrix} −0.1083−1.42780.47060.6814−1.88280.87461.17520.5529−0.5013−0.54052.1013−0.6921 −0.91540.4337−0.48591.40020.26821.05840.80230.3224−0.0986−0.93780.4463−0.6655
输出为
[ [ − 0.1538 0.5740 1.0290 − 0.5521 − 1.3698 − 1.7891 0.4555 1.8825 0.3797 0.7520 − 0.5160 − 0.6918 ] [ − 1.4311 1.7217 0.9076 − 1.4617 0.4057 0.1804 0.2542 0.4230 − 0.8464 1.2564 − 0.3189 − 1.0908 ] ] \begin{bmatrix} \begin{bmatrix} -0.1538& 0.5740& 1.0290& -0.5521\\ -1.3698& -1.7891& 0.4555& 1.8825\\ 0.3797& 0.7520& -0.5160& -0.6918\\ \end{bmatrix}& \begin{bmatrix} -1.4311& 1.7217& 0.9076& -1.4617\\ 0.4057& 0.1804& 0.2542& 0.4230\\ -0.8464& 1.2564& -0.3189& -1.0908 \end{bmatrix} \end{bmatrix} −0.1538−1.36980.37970.5740−1.78910.75201.02900.4555−0.5160−0.55211.8825−0.6918 −1.43110.4057−0.84641.72170.18041.25640.90760.2542−0.3189−1.46170.4230−1.0908

更多推荐

所有评论(0)