TensorFlow 1.12中scatter_update结合feed数据报错求助
解决TensorFlow中
scatter_update无法作用于Placeholder的问题 首先咱们得明确一个核心点:tf.scatter_update这个API只能用来更新tf.Variable对象,它依赖变量的可修改状态和专属属性来工作。而tf.placeholder只是一个接收外部输入的张量节点,本身没有可更新的内部状态,也不具备scatter_update需要的_lazy_read这类变量专属方法,这就是你遇到AttributeError的根本原因。
再看你第二次尝试的代码里还有两个小问题:
- 直接对
placeholder(也就是init_val=x)调用scatter_update,完全违背了这个API的使用前提; - 代码里调用了
tf.global_variables_initializer(),但此时图里根本没有定义任何Variable,这一步完全多余,不过不是报错的核心原因。
正确的实现思路
如果你想用外部feed的数据作为初始值,再执行scatter_update,得先把feed的值赋值给一个Variable,再对这个变量进行更新操作。具体代码如下:
import tensorflow as tf import numpy as np with tf.Session() as sess: # 定义placeholder接收外部初始值 x = tf.placeholder(tf.float32, shape=(3, 2)) # 创建Variable,用placeholder的值初始化它 init_val = tf.Variable(x, trainable=False) indices = tf.constant([0, 1]) # 对Variable执行scatter_update update = tf.scatter_update(init_val, indices, tf.ones((2, 2))) # 初始化变量时传入placeholder的feed值 sess.run(init_val.initializer, feed_dict={x: np.zeros((3, 2))}) print(sess.run(update))
代码说明
- 先通过
placeholder接收你要传入的初始数据; - 创建
tf.Variable时把placeholder作为初始值,设置trainable=False(如果不需要训练这个变量的话); - 初始化变量时,要调用
init_val.initializer并传入feed_dict,这样变量就会用你传入的np.zeros((3,2))作为初始值; - 最后调用
tf.scatter_update就完全符合API要求了,因为操作的是标准的Variable对象。
运行这段代码后,会得到和你第一个代码片段一致的输出:
[[1. 1.] [1. 1.] [0. 0.]]
内容的提问来源于stack exchange,提问作者Vincent Stimper
相关产品推荐
相关产品推荐

