NumPy ndarray如何根据指定轴的索引数组实现切片取值?
Numpy 指定维度索引取值的Pythonic实现
核心实现
最简洁的实现直接调用numpy内置的take_along_axis方法,一行代码即可完成需求:
import numpy as np # 核心取值逻辑 result = np.take_along_axis(A, Indices[:, None, :], axis=1).squeeze()
逻辑说明:
Indices[:, None, :]是给索引数组在要取值的第1轴新增一个维度,让索引数组的维度和输入数组A的维度对齐,匹配take_along_axis的参数要求squeeze()用来去掉取值后多余的第1轴维度,最终得到和Indices形状一致的3×4结果数组
如果你偏好原生高级索引写法,也可以用如下代码实现,效果完全一致:
result = A[np.arange(A.shape[0])[:, np.newaxis], Indices, np.arange(A.shape[2])]
完整验证示例
你给出的示例可以用如下代码复现,运行结果和期望完全匹配:
import numpy as np # 构造示例数组A A = np.array([ [[0.95220166, 0.49801865, 0.83217126, 0.33361628], [0.31751156, 0.85899736, 0.81965214, 0.62465746], [0.69251917, 0.83201231, 0.6089141, 0.36589825], [0.96674647, 0.6056233, 0.45515703, 0.90552863], [0.94524208, 0.42422369, 0.91633385, 0.53177495]], [[0.02883774, 0.18012477, 0.64642352, 0.21295456], [0.88475705, 0.76020851, 0.6888415, 0.47958142], [0.17306953, 0.94981064, 0.91468365, 0.37297622], [0.75924232, 0.27537972, 0.68803293, 0.0904176], [0.14596762, 0.70103752, 0.06090593, 0.07920207]], [[0.11092702, 0.58002663, 0.13553706, 0.89662211], [0.09146413, 0.86212582, 0.65908978, 0.2995175], [0.29025485, 0.60788672, 0.98595003, 0.06762369], [0.56136928, 0.09623415, 0.20178919, 0.46531331], [0.28628325, 0.28215312, 0.39670151, 0.68243605]] ]) # 构造索引数组 Indices = np.array([ [3, 1, 2, 1], [3, 2, 0, 4], [3, 3, 1, 2] ]) # 执行取值 result = np.take_along_axis(A, Indices[:, None, :], axis=1).squeeze() print(result)
运行输出:
[[0.96674647 0.85899736 0.6089141 0.62465746] [0.75924232 0.94981064 0.64642352 0.07920207] [0.56136928 0.09623415 0.65908978 0.06762369]]
内容的提问来源于stack exchange,提问作者C. Wang
相关产品推荐
相关产品推荐

