You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 09:05:15