TensorFlow中使用与不使用上下文管理器定义tf.Session为何行为不同?
嘿,最近我踩了个TensorFlow的小坑——用上下文管理器创建会话和直接实例化会话,居然会导致tf.summary.merge_all()的行为完全不一样!给你看具体的对比示例,咱们一步步说清楚为什么。
一、使用上下文管理器的示例
先看能正常工作的代码:
import tensorflow as tf graph = tf.Graph() # 在自定义图的上下文里定义变量和summary with graph.as_default(): x = tf.Variable(0) tf.summary.scalar("x", x) # 用上下文管理器创建会话 with tf.Session(graph=graph) as sess: # 在会话上下文里合并所有summary summaries = tf.summary.merge_all() print("Operations:", sess.graph.get_operations()) print("\nSummaries:", summaries)
运行结果(截取关键部分):
Operations: [<tf.Operation 'Variable/initial_value' type=Const>, <tf.Operation 'Variable' type=VariableV2>, <tf.Operation 'Variable/Assign' type=Assign>, <tf.Operation 'Variable/read' type=Identity>, <tf.Operation 'x' type=ScalarSummary>]
Summaries: Tensor("Merge/MergeSummary:0", shape=(), dtype=string)
这里merge_all()成功找到了我们在graph里定义的ScalarSummary操作,返回了有效的合并张量。
二、不使用上下文管理器的示例
再看会出问题的写法:
import tensorflow as tf graph = tf.Graph() with graph.as_default(): x = tf.Variable(0) tf.summary.scalar("x", x) # 直接创建会话,不进入上下文管理器 sess = tf.Session(graph=graph) summaries = tf.summary.merge_all() print("Operations:", sess.graph.get_operations()) print("\nSummaries:", summaries) sess.close()
运行结果会是这样:
Operations: [<tf.Operation 'Variable/initial_value' type=Const>, <tf.Operation 'Variable' type=VariableV2>, <tf.Operation 'Variable/Assign' type=Assign>, <tf.Operation 'Variable/read' type=Identity>, <tf.Operation 'x' type=ScalarSummary>]
Summaries: None
明明会话绑定的是同一个graph,为什么merge_all()会返回None?
三、差异的核心原因
这背后的关键是默认图的上下文绑定:
使用会话上下文管理器时:
- 自动管理会话生命周期(退出代码块时自动关闭会话);
- 临时将绑定的
graph设为当前默认图,同时把该会话设为默认会话。
所以在这个块里调用merge_all()时,它会去当前默认图(也就是我们自定义的graph)里查找所有summary操作,自然能找到。
直接创建会话时:
当前默认图还是TensorFlow启动时创建的全局默认图(不是我们自定义的graph)。tf.summary.merge_all()默认只会在当前默认图里查找summary操作,全局默认图里没有我们定义的x的summary,所以返回None。
四、不使用上下文管理器的解决办法
如果不想用上下文管理器,也能让代码正常工作——只需要在调用merge_all()时,显式把自定义图设为默认图:
import tensorflow as tf graph = tf.Graph() with graph.as_default(): x = tf.Variable(0) tf.summary.scalar("x", x) sess = tf.Session(graph=graph) # 显式进入自定义图的上下文 with graph.as_default(): summaries = tf.summary.merge_all() print("Operations:", sess.graph.get_operations()) print("\nSummaries:", summaries) sess.close()
这样运行就会和使用会话上下文管理器的结果完全一致啦!
内容的提问来源于stack exchange,提问作者desa

