如何用Numpy向量化高效提取列表/数组中符合条件的值?
用Numpy向量化方法高效分组百万级数据
嘿,你的问题我太懂了——百万级数据用嵌套循环确实会慢到让人崩溃,尤其是还要分组存储对应的值。下面给你一套基于Numpy的高效方案,既快又省内存,完美适配你的需求:
核心解决方案代码
import random import numpy as np # 生成测试数据(模拟百万级规模时只需修改range参数) x = [random.randrange(0, 10) for _ in range(0, 100)] y = [random.randrange(0, 10) for _ in range(0, 100)] z = [random.randrange(0, 10) for _ in range(0, 100)] # 1. 将列表转为Numpy数组(向量化操作的基础) x_arr = np.array(x) y_arr = np.array(y) z_arr = np.array(z) # 2. 对x数组排序,同时获取排序后的索引(相同值会被集中在一起) sorted_indices = np.argsort(x_arr) x_sorted = x_arr[sorted_indices] y_sorted = y_arr[sorted_indices] z_sorted = z_arr[sorted_indices] # 3. 找到分组边界:相邻元素不同的位置+数组首尾 diff = np.diff(x_sorted, prepend=-1, append=-1) group_boundaries = np.where(diff != 0)[0] # 4. 按边界分割数组,得到分组后的列表 xx_list = np.split(x_sorted, group_boundaries[1:-1]) y_list = np.split(y_sorted, group_boundaries[1:-1]) z_list = np.split(z_sorted, group_boundaries[1:-1]) # 可选:如果需要转回Python列表(根据实际需求选择) xx_list = [arr.tolist() for arr in xx_list] y_list = [arr.tolist() for arr in y_list] z_list = [arr.tolist() for arr in z_list]
方案优势解释
- 速度爆炸提升:Numpy的排序和分割操作都是底层C实现,完全避免了Python嵌套循环的巨大开销。对于百万级数据,嵌套循环可能需要几十秒甚至更久,而这套方法仅需几百毫秒(主要耗时在排序,时间复杂度为O(n log n))。
- 内存高效利用:不需要预先创建空列表再逐个
append,而是直接通过数组切片分割,内存布局更紧凑,避免了Python列表动态扩容的内存浪费。 - 结果完全匹配需求:排序后分割的方式会自动把相同
x值对应的y、z值分组到对应位置,和你预期的示例结果完全一致。
针对你补充的示例验证
用你给出的小例子测试,完全符合预期输出:
x = [0.1,0.1,1,0.1,2,1,0.1] y = [1.1,2.1,3,4,5,6,7] z = [10,11,12,13.1,14,15,16] x_arr = np.array(x) y_arr = np.array(y) z_arr = np.array(z) sorted_indices = np.argsort(x_arr) x_sorted = x_arr[sorted_indices] y_sorted = y_arr[sorted_indices] z_sorted = z_arr[sorted_indices] diff = np.diff(x_sorted, prepend=-1, append=-1) group_boundaries = np.where(diff != 0)[0] xx_list = [arr.tolist() for arr in np.split(x_sorted, group_boundaries[1:-1])] y_list = [arr.tolist() for arr in np.split(y_sorted, group_boundaries[1:-1])] z_list = [arr.tolist() for arr in np.split(z_sorted, group_boundaries[1:-1])] print(xx_list) # [[0.1, 0.1, 0.1, 0.1], [1, 1], [2]] print(y_list) # [[1.1, 2.1, 4, 7], [3, 6], [5]] print(z_list) # [[10, 11, 13.1, 16], [12, 15], [14]]
内容的提问来源于stack exchange,提问作者Abhinav Kumar
相关产品推荐
相关产品推荐

