Numpy花式索引如何自动补全轴维度 无需手动声明索引数组i
你写的a[:, j]和原代码结果不一致的核心原因是索引时的广播规则不匹配:原代码中i是形状为(3,1)的行索引数组,和形状为(3,2)的j广播后,会为每一行匹配对应当行的列索引,最终返回形状为(3,2)的结果;而a[:, j]会把j的所有元素应用到每一行上,返回的是形状为(3, 3, 2)的数组,不符合你的预期。
简洁实现方案
方案1:使用take_along_axis(最推荐)
Numpy内置的take_along_axis方法就是专门为这种逐行/逐列匹配索引的场景设计的,无需手动构造行索引数组,一行代码即可实现和原代码完全一致的效果:
import numpy as np j = np.array([[2,1], [0,2], [1,0]]) a = np.array([[1,2,3],[4,5,6],[6,7,8]]) res = a.take_along_axis(j, axis=1)
方案2:简化高级索引写法
如果你要继续用高级索引的写法,可以把行索引的生成逻辑直接写在索引语句里,省去单独声明i变量的步骤:
res = a[np.arange(len(j))[:, np.newaxis], j]
两种方案返回的结果都和原代码a[i, j]完全相同。
内容的提问来源于stack exchange,提问作者Make42
相关产品推荐
相关产品推荐

