You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.09 00:02:33