TensorFlow ScatterNdUpdate报错:需传入可变张量(如tf.Variable)
解决TensorFlow中
ScatterNdUpdate的"ref必须是可变张量"报错问题 为什么会触发这个报错?
你踩的这个坑其实很典型——tf.scatter_nd_update这个API的第一个参数ref有硬性要求:必须是可修改的tf.Variable实例。但你代码里传给它的self.embedded_chars是tf.nn.embedding_lookup的输出,这只是一个普通的计算张量(Tensor),并不是tf.Variable,TensorFlow当然不允许你对一个不可变的张量执行原地修改操作,所以直接抛出了这个错误。
怎么修改代码实现需求?
首先得明确你的真实需求:是想修改embedding_lookup得到的结果张量,还是想修改原始的嵌入矩阵self.W?我给你两种对应的解决方法:
情况1:修改embedding_lookup后的结果张量
因为普通张量不能原地修改,我们可以用tensor_scatter_nd_update(这个API是生成新张量,不需要Variable)来构造修改后的结果,实现代码如下:
with tf.device('/cpu:0'), tf.name_scope("embedding"): self.W = tf.Variable( tf.random_uniform([vocab_size, embedding_size], -1.0, 1.0), name="W" ) self.embedded_chars = tf.nn.embedding_lookup(self.W, self.input_x) # 假设你的embedded_chars形状是[batch_size, sequence_length, embedding_size] batch_size = tf.shape(self.embedded_chars)[0] # 生成要置0的所有位置的索引 indices = [] for i in range(1, sequence_length - 2): # 每个batch的第i个位置都要置0 batch_indices = tf.stack([tf.range(batch_size), tf.fill([batch_size], i)], axis=1) indices.append(batch_indices) indices = tf.concat(indices, axis=0) # 构造要更新的值:全0,每个位置对应embedding_size个元素 updates = tf.zeros([tf.shape(indices)[0], embedding_size]) # 生成修改后的张量(注意用tensor_scatter_nd_update,不是scatter_nd_update) self.embedded_chars = tf.tensor_scatter_nd_update(self.embedded_chars, indices, updates) self.embedded_chars_expanded = tf.expand_dims(self.embedded_chars, ...)
这里用tensor_scatter_nd_update直接生成修改后的新张量,完全符合你的需求,而且不需要依赖tf.Variable。
情况2:修改原始的嵌入矩阵self.W
如果你的真实需求是更新嵌入矩阵里的某些行(比如某些词汇的嵌入向量),那可以直接对self.W使用scatter_nd_update,因为它本身就是tf.Variable:
with tf.device('/cpu:0'), tf.name_scope("embedding"): self.W = tf.Variable( tf.random_uniform([vocab_size, embedding_size], -1.0, 1.0), name="W" ) # 假设你要更新的词汇索引是target_words_indices target_words_indices = [...] # 替换成你实际要更新的词汇索引列表 # 构造更新值:全0的嵌入向量 updates = tf.zeros([len(target_words_indices), embedding_size]) # 对Variable执行更新操作 updated_W = tf.scatter_nd_update(self.W, tf.expand_dims(target_words_indices, 1), updates) # 用更新后的嵌入矩阵做lookup self.embedded_chars = tf.nn.embedding_lookup(updated_W, self.input_x) self.embedded_chars_expanded = tf.expand_dims(self.embedded_chars, ...)
注意:scatter_nd_update会原地修改self.W的值,如果是在训练流程中使用,要确保这个操作的执行时机符合你的预期。
内容的提问来源于stack exchange,提问作者X. L
相关产品推荐
相关产品推荐

