如何将TF1训练的Protobuf模型加载到TF2?(遇TensorFlow bug)
解决TF2加载TF1保存的Stable-Baselines模型时的输出形状不匹配问题
嘿,我碰到过类似的问题,你这个情况确实是TensorFlow的已知bug——TF2在兼容TF1保存的带BatchNormalization梯度节点的SavedModel时,对梯度节点的形状处理出了问题。给你两个靠谱的解决办法:
方法一:重新保存模型,只保留推理图**(推荐)**
既然你加载模型是用来推理而不是继续训练,那完全没必要把训练用的梯度节点也保存下来。修改TF1里的保存代码,只导出前向推理需要的部分:
with model.graph.as_default(): # 指定推理用的输入和输出节点 inputs = {"obs": model.act_model.obs_ph} outputs = {"action": model.act_model._policy_proba} # 用SavedModelBuilder构建只含推理签名的模型 builder = tf.saved_model.builder.SavedModelBuilder('tensorflow_model_infer') signature = tf.saved_model.signature_def_utils.predict_signature_def(inputs=inputs, outputs=outputs) builder.add_meta_graph_and_variables( model.sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={ tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: signature }, strip_default_attrs=True # 移除不必要的默认属性,减少兼容性问题 ) builder.save()
然后在TF2里加载这个优化后的模型就没问题了:
import tensorflow as tf # 先确认模型合法性 assert tf.saved_model.contains_saved_model('tensorflow_model_infer') # 加载模型 model_loaded = tf.saved_model.load('tensorflow_model_infer') # 获取推理函数 infer_fn = model_loaded.signatures['serving_default'] # 测试一下(记得把输入形状换成你模型实际的输入尺寸) test_obs = tf.random.normal((1, 8, 8, 4)) # 示例形状,按需调整 result = infer_fn(obs=test_obs) print(result['action'].numpy())
方法二:临时绕过TF2的形状检查**(不推荐长期用)**
如果没法重新保存模型,可以试试在TF2里关闭Eager执行,同时调整环境变量跳过部分检查,但这个方法可能在后续TF版本里失效:
import os # 关闭确定性操作检查和冗余日志 os.environ['TF_DETERMINISTIC_OPS'] = '0' os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' import tensorflow as tf # 禁用Eager执行,兼容TF1的计算图模式 tf.compat.v1.disable_eager_execution() # 用TF1兼容模式加载模型 with tf.compat.v1.Session() as sess: model_loaded = tf.saved_model.load_v2('tensorflow_model') infer_fn = model_loaded.signatures['serving_default'] sess.run(tf.compat.v1.global_variables_initializer()) # 测试推理 test_obs = tf.random.normal((1, 8, 8, 4)) result = sess.run(infer_fn(obs=test_obs)) print(result['action'])
问题根源
你的模型里用了多个BatchNormalization层,TF1的tf.saved_model.simple_save会把整个计算图(包括训练时的梯度计算节点)都保存下来。而TF2在解析FusedBatchNormGradV3这个梯度节点时,对输出形状的处理存在bug——训练时可能存在批量大小为0的动态形状,和推理时的实际形状(比如你这里的64)不匹配,就抛出了那个错误。
所以最稳妥的还是方法一,只保存推理需要的部分,彻底避开这个兼容性bug。
内容的提问来源于stack exchange,提问作者VSZM
相关产品推荐
相关产品推荐

