tf.assign()迭代增量异常:每次增2而非1的原因及解决方法
问题分析与解决方案
嘿,这个问题我太熟悉了,咱们一步步来拆解:
1. 为什么val每次会增加2?
核心原因是你在每次循环里执行了两次assign_op操作!
先看你循环里的代码:
sess.run(assign_op) print(iter, x.eval(),y.eval(),"val",val.eval(),"assign_op", assign_op.eval())
在TensorFlow旧版的图模式里,assign_op是一个计算图节点,它的逻辑是把val+1的值重新赋值给val。而:
sess.run(assign_op)会执行一次这个赋值操作,此时val已经加1;- 后面的
assign_op.eval()本质上等价于sess.run(assign_op),又执行了一次赋值操作,val再次加1。
两次执行下来,每次循环val就会增加2,这就是你看到异常结果的原因。
2. 如何修改让val每次仅增加1?
只需要保证**每次循环只执行一次assign_op**就行,这里有两种简洁的修改方式:
方式一:移除对assign_op的重复eval
既然已经通过sess.run(assign_op)完成了赋值,就不需要再eval它了,直接打印val的当前值即可:
iters = 10 x = tf.constant(1, dtype=tf.float32, name="X") y = tf.constant(2, dtype=tf.float32, name="y") val = tf.Variable(y-x, name="val") assign_op = tf.assign(val, val+1) init = tf.global_variables_initializer() with tf.Session() as sess: sess.run(init) print("val_init",val.eval()) for iter in range(iters): sess.run(assign_op) # 只打印val的当前值,不再触碰assign_op print(iter, x.eval(), y.eval(), "val", val.eval())
方式二:一次性获取所有需要的值(更高效)
TensorFlow中,sess.run()可以同时获取多个张量的结果,这样能减少图的执行次数,效率更高:
iters = 10 x = tf.constant(1, dtype=tf.float32, name="X") y = tf.constant(2, dtype=tf.float32, name="y") val = tf.Variable(y-x, name="val") assign_op = tf.assign(val, val+1) init = tf.global_variables_initializer() with tf.Session() as sess: sess.run(init) print("val_init",val.eval()) for iter in range(iters): # 一次run获取赋值后的val、x、y的值 new_val, x_val, y_val = sess.run([assign_op, x, y]) print(iter, x_val, y_val, "val", new_val)
这里assign_op执行后会返回新的val值,所以直接用这个结果就好,不需要再单独eval val了。
运行修改后的代码,val就会每次循环只增加1啦!
内容的提问来源于stack exchange,提问作者sebtac
相关产品推荐
相关产品推荐

