为何PyTorch中BatchNorm单batch下训练与eval模式输出不同?
PyTorch BatchNorm单次batch下两种输出为啥不一样?
核心原因很直接:训练模式下的running均值/方差是指数移动平均更新,不是直接替换成当前batch的统计值。
- 刚初始化的BatchNorm,
running_mean默认是全0张量,running_var默认是全1张量。 - 训练模式(默认开启)下第一次传入batch计算时,running统计值会按加权公式更新:
running_mean = (1 - momentum)*初始均值 + momentum*当前batch均值(默认momentum=0.1)running_var = (1 - momentum)*初始方差 + momentum*当前batch方差
说白了就是旧的running值和当前batch的统计值按比例混合,不是直接用当前batch的值覆盖。 - 切换到评估模式后,BatchNorm会用更新后的running均值/方差做归一化,而非当前batch的统计结果,所以
out_2和用当前batch统计值算出的out_1自然不相等。
如果想让两者相等,可以手动把running统计值替换成当前batch的结果,再切评估模式验证:
import torch test = torch.rand((2,10)) norm = torch.nn.BatchNorm1d(10) out_1 = norm(test) # 计算当前batch的均值和方差(注意BatchNorm默认用除以batch size的方差,即unbiased=False) batch_mean = test.mean(dim=0) batch_var = test.var(dim=0, unbiased=False) # 替换running统计值 norm.running_mean.data.copy_(batch_mean) norm.running_var.data.copy_(batch_var) norm.train(False) out_2 = norm(test) print(torch.allclose(out_1, out_2)) # 会输出True(忽略浮点误差)
内容的提问来源于stack exchange,提问作者h. Oo
相关产品推荐
相关产品推荐

