NumPy中如何用一个ndarray的argmax索引取另一数组对应值
NumPy无循环取argmax对应元素的实现方法
完全可以去掉显式循环,用NumPy原生的高级索引机制一次性完成取值,运行效率比Python层循环高很多,结果和你写的循环逻辑完全一致。
你沿axis=0计算得到的a_argmax,存的是每一列最大值对应的行索引,长度等于数组的列数。这时候只需要生成一个和a_argmax等长的顺序列索引(从0到列数-1的整数序列),把两组索引同时传入数组b的索引位,就能直接批量取出所有目标位置的元素。
实现代码
import numpy as np a = np.array([[10, 20], [15, 5]]) b = np.array([[7, 6], [8, 9]]) a_argmax = a.argmax(axis=0) # 无循环一次性提取对应元素 result = b[a_argmax, np.arange(len(a_argmax))] print(result) # 输出: [8 6]
逻辑说明
对照你原来的循环逻辑看对应关系:
- 循环里列索引
i的遍历范围是0到1,正好对应np.arange(len(a_argmax))生成的[0,1] - 每个列对应的行索引
m就是a_argmax里存的[1, 0] - 高级索引会按位置配对坐标:
(1,0)、(0,1),正好取出b[1,0]=8、b[0,1]=6,和循环输出完全一致
如果后续你是沿axis=1计算argmax,要取每行最大值对应的另一个数组元素,只要调换索引顺序即可,写法为b[np.arange(b.shape[0]), a_argmax],可以适配不同轴向的计算需求。
内容的提问来源于stack exchange,提问作者Vitaly Olegovitch
相关产品推荐
相关产品推荐

