求Python中按指定区间分组二维数组并保留索引的高效方法
优化二维数组区间分组的高效方案
原代码的性能瓶颈
你当前的三重嵌套循环逐个判断元素所属区间,时间复杂度为O(MNK)(M、N是二维数组的行列数,K是区间数量),当数据规模较大时,这种方法会因为大量重复判断变得非常低效。
高效解决方案:用numpy.digitize替代循环
numpy.digitize是NumPy专为区间分组设计的函数,能一次性为所有元素匹配对应的区间索引,把时间复杂度降到O(M*N),大幅提升处理速度。
针对你提供的示例的优化实现
import numpy as np intervals = np.arange(0, 5000, 500) data = [[1091, 781, 675, 3644, 677], [91, 751, 2675, 644, 75], [3791, 4756, 3675, 644, 3669], [791, 1754, 1679, 4643, 2675], [2991, 2751, 675, 2629, 2659]] # 转换为numpy数组,方便批量处理 data_np = np.array(data) # 获取每个元素对应的区间索引,right=False对应左闭右开区间,和你的判断逻辑完全匹配 indices = np.digitize(data_np, intervals, right=False) # 初始化输出列表 output = [[] for _ in range(len(intervals))] # 遍历扁平化后的元素和索引,分配到对应分组 for idx, val in zip(indices.flatten(), data_np.flatten()): # digitize返回的索引从1开始计数,对应到output的索引需要减1 if 0 < idx <= len(intervals): output[idx-1].append(val) # 输出结果和你预期一致 print(output)
关键细节解释
np.digitize(data_np, intervals, right=False):right=False指定区间为[intervals[k], intervals[k+1]),完美匹配你代码中的data[i][j] >= intervals[k] and data[i][j] < intervals[k+1]条件。函数返回的索引idx表示元素落在第idx个区间,因此需要减1对应到output的列表索引。flatten():把二维数组转为一维,方便一次性遍历所有元素和对应的区间索引。
扩展:保留元素的原始位置索引
如果需要同时记录每个元素在原二维数组中的行、列位置,可以用以下方式实现:
# 获取所有元素的行、列索引 rows, cols = np.indices(data_np.shape) # 扁平化所有数据 flat_vals = data_np.flatten() flat_rows = rows.flatten() flat_cols = cols.flatten() indices = np.digitize(flat_vals, intervals, right=False) # 初始化带位置信息的输出 output_with_indices = [[] for _ in range(len(intervals))] for idx, val, r, c in zip(indices, flat_vals, flat_rows, flat_cols): if 0 < idx <= len(intervals): output_with_indices[idx-1].append((val, (r, c))) print(output_with_indices)
性能对比
假设处理一个1000×1000的大型数组,原三重循环可能需要数秒甚至更久,而np.digitize方案仅需几毫秒,性能提升非常显著。
内容的提问来源于stack exchange,提问作者hwerner
相关产品推荐
相关产品推荐

