如何在多个TensorFlow会话中复用已保存的计算图?
解决TensorFlow多Session导致变量未初始化的问题
我太懂这个坑了!你现在遇到的问题核心点很明确:两个独立的Session没有共享已恢复的变量状态。
你看,虽然你用graph.as_default()固定了计算图,但每个Session都是计算图的独立执行上下文——第一个Session里完成模型restore后,所有变量的初始化状态只存在于这个会话内部;第二个新Session启动时,相当于从零开始,根本没拿到之前加载的权重,所以才会报错所有变量都未初始化。
最直接的解决方案:共用同一个Session
这是最简单也最推荐的做法,把模型加载和运行操作放在同一个Session里完成,代码改成这样:
graph = tf.get_default_graph() with graph.as_default(): with tf.Session() as sess: resSaver = tf.train.import_meta_graph(lastModel) resSaver.restore(sess, checkpoint_file) x = graph.get_operation_by_name("output/x").outputs[0] # 直接在这个会话里执行你的运行逻辑 sess.run(x, feed_dict={...})
特殊场景:必须用两个Session的处理方式
如果因为某些特殊需求,你确实需要分开两个Session,那可以先在第一个Session里把所有变量的值保存下来,再在第二个Session里初始化后赋值:
graph = tf.get_default_graph() with graph.as_default(): with tf.Session() as load_sess: resSaver = tf.train.import_meta_graph(lastModel) resSaver.restore(load_sess, checkpoint_file) x = graph.get_operation_by_name("output/x").outputs[0] # 把所有全局变量的当前值存成字典 all_var_values = {var.name: load_sess.run(var) for var in tf.global_variables()} with tf.Session() as run_sess: # 先初始化所有变量 run_sess.run(tf.global_variables_initializer()) # 把之前保存的变量值逐个赋值进去 for var_name, value in all_var_values.items(): # 注意变量张量的名称要加":0"后缀 var_tensor = graph.get_tensor_by_name(f"{var_name}:0") run_sess.run(var_tensor.assign(value)) # 现在可以正常运行了 run_sess.run(x, feed_dict={...})
不过这种方法比较繁琐,除非真的有必要,否则优先用第一种方案。
总结一下:TensorFlow的Session是状态的载体,restore的变量状态只属于当前Session,跨Session不会自动共享,所以尽量让模型加载和运行在同一个会话里完成。
内容的提问来源于stack exchange,提问作者lorenzo
相关产品推荐
相关产品推荐

