【问题标题】:How to implement Numpy where index in TensorFlow?如何在 TensorFlow 中实现 Numpy where 索引?
【发布时间】:2018-07-30 07:23:06
【问题描述】:

我有以下使用numpy.where的操作:

    mat = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=np.int32)
    index = np.array([[1,0,0],[0,1,0],[0,0,1]])
    mat[np.where(index>0)] = 100
    print(mat)

如何在TensorFlow中实现等价物?

mat = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=np.int32)
index = np.array([[1, 0, 0], [0, 1, 0], [0, 0, 1]])
tf_mat = tf.constant(mat)
tf_index = tf.constant(index)
indi = tf.where(tf_index>0)
tf_mat[indi] = -1   <===== not allowed 

【问题讨论】:

  • np.where() 在这里不需要:mat[index &gt; 0] = 100
  • 不管怎样,在tensorflow中如何满足相同的意图?

标签: python numpy tensorflow


【解决方案1】:

假设您想要创建一个带有一些替换元素的新张量,而不是更新变量,您可以执行以下操作:

import numpy as np
import tensorflow as tf

mat = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=np.int32)
index = np.array([[1, 0, 0], [0, 1, 0], [0, 0, 1]])
tf_mat = tf.constant(mat)
tf_index = tf.constant(index)
tf_mat = tf.where(tf_index > 0, -tf.ones_like(tf_mat), tf_mat)
with tf.Session() as sess:
    print(sess.run(tf_mat))

输出:

[[-1  2  3]
 [ 4 -1  6]
 [ 7  8 -1]]

【讨论】:

  • 这不是我想要的!
  • @Roby Umh,你能澄清一下你到底想要什么吗?或者答案不符合您的需求怎么办?
  • @Roby 哎呀抱歉,刚刚意识到我在输出中复制了错误的矩阵(代码仍然很好......),现在已修复。
【解决方案2】:

您可以通过tf.where获取索引,然后您可以运行索引,或者使用tf.gather从原始数组中收集数据,或者使用tf.scatter_update更新原始数据,tf.scatter_nd_update用于多维更新。

mat = tf.Variable([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=tf.int32)
index = tf.Variable([[1,0,0],[0,1,0],[0,0,1]])
idx = tf.where(index>0)
tf.scatter_nd_update(mat, idx, /*values you want*/)

请注意,更新值应与 idx 的第一维大小相同。

见https://www.tensorflow.org/api_guides/python

【讨论】:

  • idx 用在scatter_update 是一维的,tf.where return 在这里不起作用!
  • 如果要使用多于一维更新,请使用tf.scatter_nd_update 而不是tf.scatter_update
  • 是的,tf.scatter_nd_update 有效,它是@jdehesa 建议的tf.where() 的替代方法。不过谢谢。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2017-11-18
  • 1970-01-01
  • 1970-01-01
  • 2017-08-25
  • 2021-04-08
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多