使用文档:https://docs.scipy.org/doc/numpy-1.13.0/user/basics.indexing.html
以下实现应该适用于某些 numpy ndarray 的任意数量的维度/形状。
首先我们需要一组多维索引和一些示例数据:
import numpy as np
y = np.arange(35).reshape(5,7)
print(y)
indexlist = [[0,1], [0,2], [3,3]]
print ('indexlist:', indexlist)
要真正提取直观的结果,诀窍是使用转置:
indexlisttranspose = np.array(indexlist).T.tolist()
print ('indexlist.T:', indexlisttranspose)
print ('y[indexlist.T]:', y[ tuple(indexlisttranspose) ])
进行以下终端输出:
y: [[ 0 1 2 3 4 5 6]
[ 7 8 9 10 11 12 13]
[14 15 16 17 18 19 20]
[21 22 23 24 25 26 27]
[28 29 30 31 32 33 34]]
indexlist: [[0, 1], [0, 2], [3, 3]]
indexlist.T: [[0, 0, 3], [1, 2, 3]]
y[indexlist.T]: [ 1 2 24]
元组...修复了我们可能会导致的未来警告:
print ('y[indexlist.T]:', y[ indexlisttranspose ])
FutureWarning: Using a non-tuple sequence for multidimensional indexing is deprecated; use `arr[tuple(seq)]` instead of `arr[seq]`.
In the future this will be interpreted as an array index,
`arr[np.array(seq)]`, which will result either in an error or a
different result.
print ('y[indexlist.T]:', y[ indexlisttranspose ])
y[indexlist.T]: [ 1 2 24]