TensorFlow V1能否部署?V1.12分文件训练测试模型加载报错问询
关于TensorFlow V1部署与跨文件加载模型的问题解答
首先明确说:TensorFlow V1完全可以用于部署,你遇到的FailedPreconditionError不是因为不能跨文件加载模型,而是加载流程和模型实例化的顺序出了问题,导致BatchNorm的移动均值这类变量没有被正确初始化或恢复。
问题根源
你在test.py里的操作顺序是:先加载训练好的meta图,再创建EnvModel实例。这会导致:
- 加载的meta图是训练时的图结构,而新创建的
EnvModel会在当前会话的图中新增一套变量(包括报错的BatchNorm_1/moving_mean_5)。 - 这些新增的变量不在你恢复的checkpoint里,也没有被初始化,所以执行
sess.run时就会触发未初始化变量的错误。
两种可行的解决方法
方法1:复用训练时的图结构,避免重新创建模型实例
既然你已经通过import_meta_graph加载了训练时的完整图,直接从图中提取需要的张量即可,不需要再新建EnvModel:
# test.py 修正后的代码 with tf.Session() as sess: # 加载训练时保存的meta图和变量 restorer = tf.train.import_meta_graph(save_path + '.meta') restorer.restore(sess, save_path) # 获取已加载图中的张量(注意名称要和训练时的张量名称完全一致) graph = tf.get_default_graph() # 示例:假设训练时EnvModel的张量命名带有前缀"env_model/" est_next_state = graph.get_tensor_by_name('env_model/est_next_state:0') loss = graph.get_tensor_by_name('env_model/loss:0') cur_state = graph.get_tensor_by_name('env_model/cur_state:0') next_state = graph.get_tensor_by_name('env_model/next_state:0') actions = graph.get_tensor_by_name('env_model/actions:0') done_flags = graph.get_tensor_by_name('env_model/done_flags:0') phase = graph.get_tensor_by_name('env_model/phase:0') # 将这些张量传入do_eval函数 est, actual, error = do_eval(sess, test_df, est_next_state, loss, cur_state, next_state, actions, done_flags, phase)
然后修改do_eval函数,直接使用传入的张量:
def do_eval(sess, test_df, est_next_state_tensor, loss_tensor, cur_state_tensor, next_state_tensor, actions_tensor, done_flags_tensor, phase_tensor): # 处理test_df得到states, next_states, actions, done_flags等数据 # ...(你的数据处理代码) est_next_state, loss = sess.run( [est_next_state_tensor, loss_tensor], feed_dict={ cur_state_tensor: states, next_state_tensor: next_states, actions_tensor: actions, done_flags_tensor: done_flags, phase_tensor: False # 测试时phase设为False,表示关闭训练模式(比如BatchNorm用移动均值) } ) # 计算est, actual, error的逻辑 # ... return est, actual, error
方法2:先创建模型实例,再恢复变量(保持EnvModel的写法)
如果你想保留EnvModel的封装,必须确保测试时创建的模型结构和训练时完全一致,并且先创建模型再加载checkpoint:
# test.py 修正后的代码 with tf.Session() as sess: # 第一步:先创建和训练时完全一样的EnvModel实例 env_model = EnvModel() # 第二步:恢复训练时保存的变量(此时Saver会匹配env_model的变量) restorer = tf.train.Saver() restorer.restore(sess, save_path) # 第三步:执行评估 est, actual, error = do_eval(sess, test_df, env_model)
注意:
- 训练时保存模型要使用
tf.train.Saver()(默认会保存所有变量,包括BatchNorm的moving_mean和moving_variance),确保checkpoint包含所有需要的变量。 EnvModel的结构(比如网络层数、BatchNorm的位置、变量命名规则)必须和训练时完全一致,否则会出现变量不匹配的错误。
额外的部署建议
如果你要把模型用于生产部署,更推荐在训练时导出SavedModel格式,这是TensorFlow官方推荐的跨环境部署格式:
# 训练时导出SavedModel tf.saved_model.simple_save( sess, export_dir='./saved_model', inputs={'cur_state': env_model.cur_state, 'actions': env_model.actions, ...}, outputs={'est_next_state': env_model.est_next_state} )
测试或部署时加载:
# test.py 加载SavedModel with tf.Session() as sess: tf.saved_model.loader.load(sess, ['serve'], './saved_model') graph = tf.get_default_graph() # 获取输入输出张量 cur_state = graph.get_tensor_by_name('cur_state:0') est_next_state = graph.get_tensor_by_name('est_next_state:0') # 执行评估逻辑
内容的提问来源于stack exchange,提问作者calveeen
相关产品推荐
相关产品推荐

