【问题标题】:TensorFlow - numpy-like tensor indexingTensorFlow - 类似 numpy 的张量索引
【发布时间】:2015-11-16 13:41:10
【问题描述】:

在 numpy 中,我们可以这样做:

x = np.random.random((10,10))
a = np.random.randint(0,10,5)
b = np.random.randint(0,10,5)
x[a,b] # gives 5 entries from x, indexed according to the corresponding entries in a and b

当我在 TensorFlow 中尝试类似的东西时:

xt = tf.constant(x)
at = tf.constant(a)
bt = tf.constant(b)
xt[at,bt]

最后一行给出了“Bad slice index tensor”异常。 TensorFlow 似乎不支持像 numpy 或 Theano 这样的索引。

有人知道是否有 TensorFlow 方法可以做到这一点(通过任意值索引张量)。我已经看过 tf.nn.embedding 部分,但我不确定它们是否可以用于此目的,即使可以,对于如此简单的事情来说,这也是一个巨大的解决方法。

(现在,我将来自x 的数据作为输入提供并在numpy 中进行索引,但我希望将x 放入TensorFlow 中以获得更高的效率)

【问题讨论】:

标签: python tensorflow


【解决方案1】:

您现在实际上可以使用tf.gather_nd 做到这一点。假设您有一个矩阵m,如下所示:

| 1 2 3 4 |
| 5 6 7 8 |

并且您想构建一个大小为 3x2 的矩阵 r,由 m 的元素构建而成,如下所示:

| 3 6 |
| 2 7 |
| 5 3 |
| 1 1 |

r 的每个元素对应m 的行和列,您可以使用这些索引拥有矩阵rowscols(从零开始,因为我们是在编程,而不是在做数学!) :

       | 0 1 |         | 2 1 |
rows = | 0 1 |  cols = | 1 2 |
       | 1 0 |         | 0 2 |
       | 0 0 |         | 0 0 |

你可以像这样堆叠成一个 3 维张量:

| | 0 2 | | 1 1 | |
| | 0 1 | | 1 2 | |
| | 1 0 | | 2 0 | |
| | 0 0 | | 0 0 | |

这样,您可以从mr 通过rowscols 得到如下:

import numpy as np
import tensorflow as tf

m = np.array([[1, 2, 3, 4], [5, 6, 7, 8]])
rows = np.array([[0, 1], [0, 1], [1, 0], [0, 0]])
cols = np.array([[2, 1], [1, 2], [0, 2], [0, 0]])

x = tf.placeholder('float32', (None, None))
idx1 = tf.placeholder('int32', (None, None))
idx2 = tf.placeholder('int32', (None, None))
result = tf.gather_nd(x, tf.stack((idx1, idx2), -1))

with tf.Session() as sess:
    r = sess.run(result, feed_dict={
        x: m,
        idx1: rows,
        idx2: cols,
    })
print(r)

输出:

[[ 3.  6.]
 [ 2.  7.]
 [ 5.  3.]
 [ 1.  1.]]

【讨论】:

  • @Mr_and_Mrs_D tf.gather_nd 的确切规格有点复杂,您可以在文档中查看。但基本上,我想要一个矩阵result,比如MxN,每个元素都取自另一个矩阵x。对于每个元素,我都有x的相应行和列;这些是idx1idx2。我将这两个堆叠起来得到一个 MxNx2 张量,我们称之为idx12tf.gather_nd 使用 idx12 的最后一个维度(大小为 2,类似于 x 中的维度数)来创建用于查找进入 result 的元素的二维索引。
  • 请将这些添加到您的答案中,我会投票赞成 - 文档,呃,有点缺乏。您仍然应该解释 MxN 与 idx1/2 的关系
  • @Mr_and_Mrs_D 我已经用更多的上下文和解释更新了答案。我希望现在更清楚了。
【解决方案2】:

LDGN 的评论是正确的。目前这是不可能的,并且是一个请求的功能。如果您关注issue#206 on github,您将在可用时获得更新。很多人都喜欢这个功能。

【讨论】:

    【解决方案3】:

    对于Tensorflow 0.11,已实现基本索引。仍然缺少更高级的索引(如布尔索引),但显然计划用于未来的版本。

    可以使用https://github.com/tensorflow/tensorflow/issues/4638跟踪高级索引

    【讨论】:

      猜你喜欢
      • 2016-03-03
      • 1970-01-01
      • 2017-05-21
      • 1970-01-01
      • 2020-12-15
      • 1970-01-01
      • 2017-08-25
      • 1970-01-01
      相关资源
      最近更新 更多