Numpy如何根据索引列表提取多维数组的对应值
Numpy按行提取指定列元素的最优实现
直接使用Numpy的高级整数索引即可实现,不需要额外循环,底层由C语言实现,是性能最优的方案。
实现代码
import numpy as np # 原始二维数组 arr = np.array([[1, 0, 0], [0, 0, 1]]) # 每行对应的待提取列索引 col_index = np.array([0, 2]) # 构造对应行索引数组,长度和行数一致 row_index = np.arange(arr.shape[0]) # 配对索引提取元素 res = arr[row_index, col_index]
运行后得到的res为array([1, 1]),完全符合需求。
原理解释
当传入两个形状相同的整数数组作为Numpy数组的两个维度索引时,Numpy会自动逐位置配对两个索引数组的取值,依次获取arr[row_index[i], col_index[i]]的元素,最终返回和索引数组形状一致的结果数组,全程无Python层的遍历开销。
避坑提示:不推荐使用
np.diag(arr[:, col_index])这类取巧写法,该方法会先生成一个大小为「行数×索引长度」的中间数组,当数据规模较大时,内存和时间开销都远高于高级索引方案。
内容的提问来源于stack exchange,提问作者rv123
相关产品推荐
相关产品推荐

