如何用纯NumPy方法从3D数组中按2D索引提取对应2D元素?
纯NumPy方式提取3D数组中对应2D索引的元素
你可以利用NumPy的整数数组索引特性,直接拆分索引的行和列维度来实现纯NumPy的高效提取,无需列表推导:
实现代码
import numpy as np input_matrix = np.array( [[[0.0, 1.5], [3.0, 3.0]], [[7.0, 5.2], [6.0, 7.0]]] ) indices = np.array([[1, 0], [1, 1]]) # 将二维索引拆分为行、列两个一维数组 rows = indices[:, 0] cols = indices[:, 1] # 直接通过数组索引提取对应元素 selected_elements = input_matrix[rows, cols]
原理说明
input_matrix的形状是(2, 2, 2),前两个维度对应[row, column]的矩阵结构,第三个维度是每个位置的[x,y]坐标。当你用rows(一维数组[1,1])和cols(一维数组[0,1])作为索引时,NumPy会将两个数组的对应位置组合成(1,0)、(1,1)这样的二维索引,直接定位到三维数组的前两个维度,最终返回每个位置的完整坐标数组。
运行后selected_elements的结果为:
array([[7. , 5.2], [6. , 7. ]])
和你原列表推导的结果完全一致,且性能更优(避免了Python层面的循环)。
关于numpy.take的补充
你之前尝试的take函数默认是基于数组扁平化后的一维索引进行提取,若要用于这个场景,需要手动计算每个二维索引对应的扁平化位置(比如row * 列数 + col),但这种方式不如直接用数组索引直观且易维护。
内容的提问来源于stack exchange,提问作者etien
相关产品推荐
相关产品推荐

