【问题标题】:How to index multidimensional array multiple times in a vectorized way numpy?如何以向量化方式numpy多次索引多维数组?
【发布时间】:2021-03-20 10:46:28
【问题描述】:

我正在尝试在 numpy.h 中索引一个多维数组(4 维)。数组的形状为 (125,125,125,3)。我有 3 个独立的 2D 索引数组列表。这些列表的大小分别为 (N,4)、(M,4) 和 (1,4)。这 3 个单独的列表代表我试图索引的 4D 数组中的行、列和深度值。例如考虑以下内容:

ix = [[0,1,2,3],
     [3,4,5,6]]
iy = [[2,3,4,5],
     [5,6,7,8]]
iz = [[1,2,3,4]]

weights.shape = (125,125,125,3)

我想用ixiyiz 中的行、列和深度索引数组的所有可能组合来索引weights。例如,如果我在每个索引矩阵中取第一行,这意味着我想在weights 中选择行[0,1,2,3]、列[2,3,4,5] 和深度值[1,2,3,4]。我总是想选择weights 的第四维中的所有元素。这意味着我实际上是在选择weights(4,4,4,3) 切片。

现在,我已经使用以下代码通过循环索引来实现这一点

w = np.empty(shape=(X,Y,Z,4,4,4,weights.ndim-1))
for i in range(X):
    for j in range(Y):
        w_ij = np.ix_(ix[i,:], iy[j,:], iz[0,:])
        w[i,j,0,:,:,:,:] = weights[w_ij[0], w_ij[1], w_ij[2], :]

我的最终目标是尽可能快地构造形状为 (N,M,1,4,4,4,3) 的多维数组w。这部分代码将运行多次,所以如果有一种使用内置 numpy 函数的矢量化方式来执行此操作,那将是理想的。

如果有任何需要澄清的问题,请告诉我。这是我第一次问关于堆栈溢出的问题,所以如果有任何不清楚或令人困惑的地方,我深表歉意!

【问题讨论】:

  • 虽然ix_ 只接受一维数组(我认为),但请查看生成的w_ij。我可以想象将这些数组泛化为同时使用所有ij。形状为 (N,4,M,4,...) 的 w 可能最容易生成,但可以稍后转置。我会在 ipython 会话上进行实验以帮助了解更多细节。
  • 谢谢@hpaulj!我相信ix_ 只接受一维数组是对的。您能否详细说明您认为我们如何将w_ij 概括为同时对所有ij 进行操作?对我来说,w 的形状是否有问题并不重要,因为我总是可以在之后修复它。如果您的建议可行,那将大大加快我的代码速度。
  • @ananda 知道了。

标签: python numpy numpy-ndarray numpy-indexing


【解决方案1】:

您可以将索引与广播结合使用来实现此目的。

import numpy as np

weights = np.random.rand(125, 125, 125, 3)

ix = np.array([[0,1,2,3], [3,4,5,6]])
iy = np.array([[2,3,4,5], [5,6,7,8]])
iz = np.array([[1,2,3,4]])

X = len(ix)
Y = len(iy)
Z = len(iz)

def compute1(weights):
    w = np.empty(shape=(X, Y, Z, 4, 4, 4, weights.ndim-1))
    for i in range(X):
        for j in range(Y):
            w_ij = np.ix_(ix[i,:], iy[j,:], iz[0,:])
            w[i,j,0,:,:,:,:] = weights[w_ij[0], w_ij[1], w_ij[2], :]
    return w

def compute2(weights):
    return weights[ix[:, None, None, :, None, None], iy[None, :, None, None, :, None], iz[None, None, :, None, None, :]]

print(np.allclose(compute1(weights), compute2(weights)))

True

基准测试-

%timeit compute1(weights)
%timeit compute2(weights)

给 -

36.7 µs ± 897 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
6.28 µs ± 62.8 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)

如您所见,对于这种大小的数据,广播解决​​方案的速度提高了大约 6 倍。

【讨论】:

  • 谢谢,这太棒了!非常感谢您的帮助。
猜你喜欢
  • 2015-04-19
  • 1970-01-01
  • 1970-01-01
  • 2016-12-07
  • 1970-01-01
  • 1970-01-01
  • 2018-03-23
  • 2020-03-25
  • 2018-01-22
相关资源
最近更新 更多