为什么Keras BatchNorm与PyTorch BatchNorm的输出结果存在明显差异?
环境版本
- Torch版本:
1.9.0+cu111 - Tensorflow-gpu版本:
2.5.0
问题描述
使用TensorFlow 2.5的BatchNormal层与PyTorch 1.9的BatchNorm2d层对相同的Tensor执行计算时,输出结果差异极大(TensorFlow输出接近1,PyTorch输出接近0)。调整momentum和epsilon参数为一致后,输出仍然存在差异。
复现代码
from torch import nn import torch x = torch.ones((20, 100, 35, 45)) a = nn.Sequential( # nn.Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), padding=0, bias=True), nn.BatchNorm2d(100) ) b = a(x) import tensorflow as tf import tensorflow.keras as keras from tensorflow.keras.layers import * x = tf.ones((20, 35, 45, 100)) a = keras.models.Sequential([ # Conv2D(128, (1, 1), (1, 1), padding='same', use_bias=True), BatchNormalization() ]) b = a(x)
原始输出结果


差异根因
- 运行模式默认值不同:PyTorch的
nn.BatchNorm2d实例默认处于训练模式,训练模式下使用当前batch的均值、方差做归一化;Keras的BatchNormalization层默认调用时处于推理模式,推理模式下使用全局滑动平均的均值、方差做归一化。
本次输入为全1张量,训练模式下当前batch均值为1、方差为0,归一化后结果为0,对应PyTorch输出;推理模式下滑动平均初始均值为0、方差为1,归一化后结果接近1,对应TensorFlow输出,正好放大了模式差异。 - 动量更新方向相反:PyTorch滑动平均更新公式为
running_mean = running_mean * momentum + current_mean * (1 - momentum),默认momentum=0.1;TensorFlow滑动平均更新公式为running_mean = running_mean * (1 - momentum) + current_mean * momentum,默认momentum=0.99,两者momentum参数需要设为互补值才能对齐。 - epsilon默认值不同:PyTorch默认
eps=1e-5,TensorFlow默认eps=1e-3,需要手动对齐。
对齐方案
训练模式对齐代码
# PyTorch 训练模式 import torch from torch import nn torch.manual_seed(42) x_torch = torch.ones((20, 100, 35, 45)) bn_torch = nn.BatchNorm2d(100, eps=1e-3, momentum=0.1) # eps对齐TF,momentum按PyTorch逻辑设0.1对应TF的0.9 bn_torch.train() # 显式指定训练模式 out_torch = bn_torch(x_torch) # TensorFlow 训练模式 import tensorflow as tf import tensorflow.keras as keras from tensorflow.keras.layers import * tf.random.set_seed(42) x_tf = tf.ones((20, 35, 45, 100)) bn_tf = BatchNormalization(epsilon=1e-3, momentum=0.9) # 动量设为0.9和PyTorch的0.1对应 out_tf = bn_tf(x_tf, training=True) # 显式指定训练模式 # 转换维度后验证差异 out_tf_np = tf.transpose(out_tf, [0,3,1,2]).numpy() out_torch_np = out_torch.detach().numpy() print("最大差异:", abs(out_tf_np - out_torch_np).max()) # 误差在1e-6量级,基本一致
推理模式对齐代码
# PyTorch 推理模式 import torch from torch import nn torch.manual_seed(42) x_torch = torch.ones((20, 100, 35, 45)) bn_torch = nn.BatchNorm2d(100, eps=1e-3, momentum=0.1) bn_torch.eval() # 显式指定推理模式 out_torch = bn_torch(x_torch) # TensorFlow 推理模式 import tensorflow as tf import tensorflow.keras as keras from tensorflow.keras.layers import * tf.random.set_seed(42) x_tf = tf.ones((20, 35, 45, 100)) bn_tf = BatchNormalization(epsilon=1e-3, momentum=0.9) out_tf = bn_tf(x_tf, training=False) # 显式指定推理模式 # 转换维度后验证差异 out_tf_np = tf.transpose(out_tf, [0,3,1,2]).numpy() out_torch_np = out_torch.detach().numpy() print("最大差异:", abs(out_tf_np - out_torch_np).max()) # 误差在1e-6量级,基本一致
内容的提问来源于stack exchange,提问作者call_me_ye
相关产品推荐
相关产品推荐

