Numpy二维数组按行选取指定索引元素的问题及解决
问题分析与解决方法
错误原因
当执行sizes[:, threshold_indeces]时,NumPy的广播机制会将形状为(N,)的threshold_indeces数组作为列索引,对每一行都提取这N个索引对应的元素,最终生成N×N的二维数组。在百万行场景下,该数组包含10¹²个元素,远超常规内存的承载上限,因此触发内存分配错误。
正确提取一维结果的方法
方法1:行索引+列索引配对
利用np.arange生成每行的索引,与threshold_indeces一一对应,直接定位目标元素:
import numpy as np sizes = np.array([[1, 3, 6, 6, 6, 7, 8, 8, 10, 10], [2, 3, 3, 7, 7, 7, 9, 9, 10, 11], [2, 3, 3, 5, 5, 6, 9, 9, 10, 11], [2, 3, 3, 9, 9, 9, 9, 9, 10, 11]]) threshold_indeces = np.argmax(sizes >= 5, axis=1) row_indices = np.arange(sizes.shape[0]) values = sizes[row_indices, threshold_indeces] print(values) # 输出:[6 7 5 9]
方法2:使用np.take_along_axis
该函数专门用于沿指定轴提取对应索引的元素,无需手动生成行索引:
values = np.take_along_axis(sizes, threshold_indeces[:, np.newaxis], axis=1).flatten() print(values) # 输出:[6 7 5 9]
优化:针对有序数组的高效索引
由于数组沿第二轴单调递增,使用np.searchsorted比argmax更高效(时间复杂度更低):
threshold = 5 threshold_indeces = np.searchsorted(sizes, threshold, side='left', axis=1) values = sizes[np.arange(sizes.shape[0]), threshold_indeces] print(values) # 输出:[6 7 5 9]
内容的提问来源于stack exchange,提问作者Edin
相关产品推荐
相关产品推荐

