如何在tf.while_loop中使用tf.scatter_update?遇属性错误求助
解决tf.while_loop中tf.scatter_update的AttributeError问题
咱们先拆解下你遇到的错误根源:AttributeError: 'Tensor' object has no attribute '_lazy_read'是因为第一次循环执行后,你把var替换成了tf.scatter_update返回的Tensor对象。而tf.scatter_update是专门用来修改Variable的操作,它的第一个参数必须是tf.Variable类型,第二次循环时传入的var已经变成了Tensor,自然就触发了错误。
至于tf.add能正常运行,是因为它只需要输入张量(Variable本身也属于张量的一种),返回的Tensor可以直接在循环里传递做判断,不需要依赖Variable的特殊属性。
那怎么解决呢?核心是要让var始终保持tf.Variable的身份,不要把它替换成更新后的Tensor。具体做法是单独定义更新操作,用tf.control_dependencies确保更新完成后再返回原变量。修改后的代码如下:
import tensorflow as tf def func(var1, cons): var1, _ = tf.while_loop(cond, body, [var1, cons], return_same_structure=True) with tf.control_dependencies([var1]): return var1 def cond(var, cons): return tf.reduce_all(tf.less(var, cons)) def body(var, cons): # 定义更新操作,但不把var重新赋值为返回值 update_op = tf.scatter_update(var, [0], var[0] + 1.0) # 确保更新操作执行完成后,再返回原Variable对象 with tf.control_dependencies([update_op]): return (var, cons) with tf.Session() as sess: x = tf.constant([10.0]) m = tf.Variable([2.0]) b = func(m, x) # 推荐用global_variables_initializer替代废弃的initialize_all_variables init = tf.global_variables_initializer() sess.run(init) print(sess.run(b))
额外补充两个小细节:
tf.scatter_update本身会原地更新Variable,它返回的Tensor只是更新后的值的快照,完全不需要把Variable替换成这个快照。tf.control_dependencies的作用是给TensorFlow明确执行顺序:必须先跑完update_op,才能执行后面的返回操作,保证每次循环都确实完成了元素更新。
内容的提问来源于stack exchange,提问作者Mehrdad
相关产品推荐
相关产品推荐

