【问题标题】:Complex indexing of a multidimensional array with indices of lower dimensional arrays in PythonPython中具有低维数组索引的多维数组的复杂索引
【发布时间】:2021-11-10 11:07:10
【问题描述】:

问题:

  • 我有一个 4 维的 numpy 数组:

    x = np.arange(1000).reshape(5, 10, 10, 2 )
    

    如果我们打印它:

  • 我想在第二个轴上找到数组的 6 个最大值的索引,但仅在最后一个轴上的第 0 个元素(图中的红色圆圈图片):

    indLargest2ndAxis = np.argpartition(x[...,0], 10-6, axis=2)[...,10-6:]
    

    这些索引的形状符合预期的 (5,10,6)。

  • 我想获得 第二个轴上这些索引的数组值,但现在是最后一个轴上的第一个元素(图像中的黄色圆圈)。它们的形状应为 (5,10,6)。如果没有矢量化,这可以通过以下方式完成:

    np.array([ [ [ x[i, j, k, 1] for k in indLargest2ndAxis[i,j]] for j in range(10) ] for i in range(5) ])
    

    但是,我想实现它的矢量化。我尝试使用以下索引进行索引:

    x[indLargest2ndAxis, 1]
    

    但我得到IndexError: index 5 is out of bounds for axis 0 with size 5。 如何以矢量化方式管理此索引组合?

【问题讨论】:

  • np.argpartition "返回一个与a 形状相同的索引数组,以分区顺序沿给定轴索引数据。"所以第一次天真的尝试是使用省略号x[..., indLargest2ndAxis, 1]。也就是说,我不太清楚“显示最后一个轴中第一个元素的值 [...]”是什么意思。它只是另一片吗?在这种情况下,花式索引(高级索引)。
  • 另请注意np.argpartition(x[..., 0], 6, ...) 可能无法达到您的预期;它将产生对x[...,0] 进行分区的索引,这样 - 沿着索引为 2 的轴 - x[..., 6, 0] 处于其排序位置,它之前的元素较小,之后的元素较大。如果 x[...,6,0] 确实是数组的(唯一的)第 6 大值,您可以通过 x[..., 6:, 0] 检索 6 个最大的元素;否则,您可能无法始终获得 6 个最大值。
  • @FirefoxMetzger 啊,我把最大和最小搞混了,我想现在indLargest2ndAxis 应该包含x[...,0] 沿第二轴的6 个最大值的索引。
  • @FirefoxMetzger 是的,我的问题是如何进行这种高级索引。 indLargest2ndAxis 保存 x 的 0、1 和 2 轴的索引,而第 3 轴应该是 1。我不明白为什么直接写x[indLargest2ndAxis, 1] 不起作用,但x[..., indLargest2ndAxis, 1] 会产生结果。当你写... 时会发生什么?而且,这个表达式并没有给出我想要的结果:它的形状为 (5, 10, 5, 10, 6),但我想要 (5, 10, 6),其值由图片。

标签: python numpy multidimensional-array indexing numpy-slicing


【解决方案1】:

啊,我想我现在得到了你想要的东西。详细的花式索引是documented here。但请注意 - 总的来说 - 这是相当沉重的东西。简而言之,花式索引允许您从源数组中获取元素(根据某些idx)并将它们放入新数组中(花式索引总是返回一个副本):

source = np.array([10.5, 21, 42])
idx = np.array([0, 1, 2, 1, 1, 1, 2, 1, 0])

# this is fancy indexing
target = source[idx]

expected = np.array([10.5, 21, 42, 21, 21, 21, 42, 21, 10.5])
assert np.allclose(target, expected)

这样做的好处是您可以使用索引数组的形状来控制结果数组的形状:

source = np.array([10.5, 21, 42])
idx = np.array([[0, 1], [1, 2]])

target = source[idx]

expected = np.array([[10.5, 21], [21, 42]])
assert np.allclose(target, expected)
assert target.shape == (2,2)

如果source 有多个维度,事情会变得更有趣。在这种情况下,您需要指定每个轴的索引,以便 numpy 知道要采用哪些元素:

source = np.arange(4).reshape(2,2)
idxA = np.array([0, 1])
idxB = np.array([0, 1])

# this will take (0,0) and (1,1)
target = source[idxA, idxB]

expected = np.array([0, 3])
assert np.allclose(target, expected)

再次注意,target 的形状与所用索引的形状相匹配。花式索引的厉害之处在于,如果需要,索引形状会被广播:

source = np.arange(4).reshape(2,2)
idxA = np.array([0, 0, 1, 1]).reshape((4,1))
idxB = np.array([0, 1]).reshape((1,2))

target = source[idxA, idxB]

expected = np.array([[0, 1],[0, 1],[2, 3],[2, 3]])
assert np.allclose(target, expected)

此时,您可以了解您的异常来自何处。你的source.ndim 是4;但是,您尝试使用 2 元组 (indLargest2ndAxis, 1) 对其进行索引。当您尝试使用indLargest2ndAxis 索引第一个轴,使用1 索引第二个轴和使用: 的所有其他轴时,Numpy 会解释这一点。显然,这是行不通的。 indLargest2ndAxis 的所有值都必须介于 0 和 4(含)之间,因为它们必须引用沿 x 第一轴的位置。

我对@9​​87654339@ 的建议是告诉numpy 你希望索引x 的最后两个轴,即你希望使用indLargest2ndAxis 索引第三轴,使用1 索引第四轴,和: 其他任何东西。

这将产生一个结果,因为indLargest2ndAxis 的所有元素都在[0, 10) 中,但会产生(5, 10, 5, 10, 6) 的形状(这不是您想要的)。有点手波,形状的第一部分(5, 10) 来自省略号(...),又名。选择一切,中间部分(5, 10, 6)来自indLargest2ndAxis根据indLargest2ndAxis的形状沿x的第三轴选择元素,最后一部分(你看不到,因为它被挤压)来自沿第四轴选择索引1。


继续您的实际问题,您可以完全避开花哨的索引项目符号并执行以下操作:

x = np.arange(1000).reshape(5, 10, 10, 2)
order = x[..., 0]
values = x[..., 1]
idx = np.argpartition(order, 4)[..., 4:]
result = np.take_along_axis(values, idx, axis=-1)

编辑:当然,你也可以使用花哨的索引;然而,它更神秘,不能很好地缩放到不同的形状:

x = np.arange(1000).reshape(5, 10, 10, 2)
indLargest2ndAxis = np.argpartition(x[..., 0], 4)[..., 4:]
result = x[np.arange(5)[:, None, None], np.arange(10)[None, :, None], indLargest2ndAxis, 1]

【讨论】:

  • 惊人的答案,非常感谢!我读过一些关于高级切片的文章,但我无法完全理解它。现在我想我终于成功了。我只是建议了一个我认为更容易理解的小重新排序。
  • 另外,如果它可能对某人有所帮助,我发现在这个问题中,高级索引在我的计算机上的速度要快 6.0 µs,而 np.take_along_axis 需要 9.6 µs。
  • @Puco4 感谢您的称赞。如果回答了问题,请记住将答案标记为正确。
猜你喜欢
  • 1970-01-01
  • 2017-07-17
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2013-11-22
  • 1970-01-01
相关资源
最近更新 更多