TensorFlow 2.3.4下Keras模型保存加载后predict返回全NaN求助
TensorFlow 2.3.4 加载SavedModel后预测全NaN故障排查与修复
故障根因
该问题是TensorFlow 2.3.x版本的已知缺陷,触发逻辑如下:
- 模型中包含
BatchNormalization层时,2.3版本的SavedModel序列化逻辑不会正确持久化BatchNorm层的滑动均值(moving_mean)、滑动方差(moving_variance)两类不可训练参数,加载后这两个参数会被初始化为无效值,推理阶段计算时出现除0、无效值传播,最终输出全NaN。你之前排查时核对的通常为模型可训练参数,未覆盖这部分不可训练参数,因此会误判为参数完全一致。 - 若训练时开启了混合精度训练(
tf.keras.mixed_precision),2.3版本对混合精度模型的保存逻辑存在漏洞,参数的浮点精度映射在序列化时丢失,加载后参数精度不匹配也会产生无效值。
修复方案
方案1:升级TensorFlow版本(最优解)
直接升级到TensorFlow 2.6及以上版本,该版本已彻底修复SavedModel序列化的相关缺陷,保存和加载逻辑完全兼容BatchNorm、混合精度训练场景。
方案2:2.3.4版本兼容修复
如果无法升级版本,可选择以下任意一种操作:
- 改用HDF5格式保存模型:
# 原保存脚本修改为 model.save('./model.h5') # 原加载脚本修改为 model = tf.keras.models.load_model('./model.h5')
HDF5格式在2.3版本中可正确持久化所有参数,不会出现不可训练参数丢失的问题。
- 若必须使用SavedModel格式,加载后手动重置BatchNorm层状态:
# 加载模型后执行 for layer in model.layers: if isinstance(layer, tf.keras.layers.BatchNormalization): layer.trainable = False # 重新编译模型,无需再次训练 model.compile()
验证标准
修复后执行推理,输出结果和训练脚本中的预测结果误差小于1e-5即可判定为修复成功。
内容的提问来源于stack exchange,提问作者Jivan
相关产品推荐
相关产品推荐

