如何根据指定行列数组从Numpy数组中提取对应元素?
从Numpy数组中按指定行列索引提取元素
给定以下Numpy数组:
import numpy as np mat = np.array([ [1, 2, 3], [4, 5, 6], [7, 8, 9] ]) row_col = np.array([ [1, 1], # 对应元素5 [2, 2], # 对应元素9 [0, 2], # 对应元素3 [1, 0] # 对应元素4 ])
要提取得到结果[5, 9, 3, 4],可以用以下两种常用方法:
方法1:拆分索引直接花式提取
这是最简洁高效的方式,把row_col拆分成独立的行、列索引数组,再通过Numpy的花式索引直接定位元素:
# 拆分出行索引和列索引数组 rows = row_col[:, 0] cols = row_col[:, 1] # 提取对应位置的元素 result = mat[rows, cols] print(result) # 输出: [5 9 3 4]
原理是Numpy支持同时传入行、列索引数组,每个索引对(rows[i], cols[i])会精准定位到mat中的对应元素,最终返回按顺序排列的结果数组。
方法2:使用np.take_along_axis提取
如果需要基于轴方向进行索引提取,可以用take_along_axis方法,步骤稍繁琐但逻辑清晰:
# 为索引添加新维度,适配轴操作要求 row_indices = row_col[:, 0, np.newaxis] col_indices = row_col[:, 1, np.newaxis] # 先按行索引提取对应行,再按列索引提取对应元素 temp = np.take_along_axis(mat, row_indices, axis=0) result = np.take_along_axis(temp, col_indices, axis=1).flatten() print(result) # 输出: [5 9 3 4]
最后通过flatten()将二维结果转为一维数组,得到目标输出。
内容的提问来源于stack exchange,提问作者TYL
相关产品推荐
相关产品推荐

