PyTorch中nn.BatchNorm2d的running_mean/running_var含义说明
nn.BatchNorm2d 中 running_mean、running_var 属性说明
核心本质
这两个属性是BatchNorm层专门为推理阶段准备的统计值,在训练过程中通过指数滑动平均累计得到,代表模型学到的全量训练数据对应通道的特征均值、方差,不是当前输入batch的实时计算结果。
关键细节
- 维度匹配:两个张量的形状均为
(num_features,),和初始化nn.BatchNorm2d时传入的通道数参数一一对应,每个输出通道单独存一组统计值。比如给BN设16个输出通道,这俩张量长度就是16。 - 训练阶段更新规则:当模型处于
train()模式时,每跑一次前向传播,层会自动按以下逻辑更新两个属性,完全不需要手动写代码干预:
这里的# momentum是BatchNorm初始化时的入参,默认值0.1 running_mean = (1 - momentum) * running_mean + momentum * batch_mean_current running_var = (1 - momentum) * running_var + momentum * batch_var_currentbatch_mean_current、batch_var_current是当前输入batch在对应通道上算出来的实时均值、方差。要特别注意:训练阶段做归一化计算时,用的就是这俩实时算出来的batch统计量,根本不会调用累计的running值。 - 推理阶段作用:当你把模型切到
eval()模式做推理时,BatchNorm会直接停掉当前batch统计量的计算,全程用训练阶段累计好的running_mean和running_var做归一化。这么设计的目的是保证推理时不管batch size是多少——哪怕一次只输入1张图——归一化的基准都是固定的,输出结果不会随输入batch的变化乱飘。 - 常见误区:
running_mean默认初始值是全0,running_var默认初始值是全1- 两个值不是所有训练batch统计量的简单算术平均,越靠后训练的batch对最终值的影响权重越高;如果训练步数太少、或者
momentum参数设得不合理,两个值和真实全量数据的统计量会有明显偏差 - 如果初始化BatchNorm时把
track_running_stats设成了False,这俩属性不会随训练更新,永远停在初始值,这种模式下训练和推理都会直接用当前batch的实时统计量做归一化
对你提供示例代码的说明
你贴的代码逻辑很明确:分别提取conv3、conv5后BN层的running_mean、running_var的均值和标准差,拼成一个8维的特征向量,这类操作一般用在模型训练状态分析、跨域数据偏移检测这类任务里。
内容的提问来源于stack exchange,提问作者yjyjy131
相关产品推荐
相关产品推荐

