TensorFlow中嵌套控制依赖上下文及测试代码执行分析
嘿,我来帮你拆解TensorFlow里嵌套控制依赖的工作机制,还有这段测试代码的执行逻辑与输出结果~
TensorFlow嵌套控制依赖上下文的工作机制
首先得搞懂tf.control_dependencies()的基本作用:它是用来强制指定的操作必须在上下文内的所有操作执行前完成的工具。当你进入一个with tf.control_dependencies([依赖操作列表])块时,TensorFlow会把这个列表里的操作加入当前上下文的控制依赖集合中——块里定义的任何操作,都会自动带上这些依赖,也就是说,必须等所有依赖操作都执行完,才能执行块内的操作。
那嵌套的情况呢?简单说就是依赖会逐层累积:内层的with块会继承外层所有已有的控制依赖,再加上自己指定的新依赖。不过要注意:如果同一个操作被多次添加为依赖(比如代码里的g_optim出现两次),TensorFlow的执行引擎会自动去重,不会重复执行同一个操作——因为它识别到这是同一个操作对象,只会执行一次。
另外要强调:控制依赖只保证依赖操作在目标操作之前执行,但如果依赖操作之间没有数据依赖,TensorFlow可能会并行执行它们;但如果像这个例子里,多个依赖操作都修改同一个变量,那数据依赖会强制它们按顺序执行(先执行修改变量的操作,再执行依赖该变量值的操作)。
测试代码的执行逻辑与输出结果
先看测试代码(注意这是TensorFlow 1.x的API,因为用了tf.Session和旧版的tf.get_variable):
from unittest import TestCase import tensorflow as tf class TestControl(TestCase): def test_control_dep(self): print(tf.__version__) a = tf.get_variable('a', initializer=tf.constant(0.0)) d_optim = tf.assign(a, a + 2) g_optim = tf.assign(a, a * 2) with tf.control_dependencies([d_optim]): with tf.control_dependencies([g_optim]): with tf.control_dependencies([g_optim]): op = tf.Print(a, [a]) with tf.Session() as sess: sess.run(tf.global_variables_initializer()) sess.run(op)
执行逻辑拆解
初始化与操作定义:
- 变量
a初始值为0.0。 d_optim是赋值操作:a = a + 2,执行后a会变为2.0。g_optim是赋值操作:a = a * 2,执行后a会变为当前值的两倍。
- 变量
嵌套依赖的最终效果:
- 最外层依赖
d_optim,中间层添加g_optim,最内层又加了一次g_optim——但重复的g_optim会被去重,所以op的最终依赖是{d_optim, g_optim}。 - 由于
g_optim依赖a的值,而d_optim会修改a,所以TensorFlow会保证d_optim先执行,再执行g_optim(数据依赖强制顺序)。
- 最外层依赖
实际执行流程:
- 调用
sess.run(op)时,先触发所有依赖操作:- 执行
d_optim:a从0.0变为0.0 + 2 = 2.0。 - 执行
g_optim:a从2.0变为2.0 * 2 = 4.0。
- 执行
- 最后执行
tf.Print操作,打印当前a的值。
- 调用
输出结果
- 首先会打印你的TensorFlow版本(比如
1.15.0之类的1.x版本)。 - 然后会输出
[4.],这是tf.Print打印的a的最终值。
内容的提问来源于stack exchange,提问作者Mr_and_Mrs_D
相关产品推荐
相关产品推荐

