如何快速从NumPy数组中按分组筛选多元素?
优化numpy按掩码分组筛选的性能
我开发了一个小型ML库,需要根据样本所属的原型,通过另一个数组的掩码从数组中筛选元素并分组,当前使用的代码如下:
for neighbor_id in np.unique(nearest_neighbors): samples = data[nearest_neighbors == neighbor_id ] # some function
这段代码的运行时间占总运行时间的95%,性能瓶颈极为突出。
8年前曾出现过类似的numpy记录数组索引过慢的问题,但我无法将当时的np.take()解决方案应用到当前场景,想知道有没有更新的优化方案?
补充基准测试
以下是针对两种方案的小型基准测试:
import numpy as np from timeit import timeit # 创建测试数据 rng = np.random.default_rng(seed=42) neighbors = rng.integers(0, 200, 100000) data = rng.random(size=(100000, 800)) def my_solution(): for neighbor_id in np.unique(neighbors): samples = data[neighbors == neighbor_id] def jeromes_solution(): index = np.argsort(neighbors) groups, offsets = np.unique(neighbors[index], return_index=True) for i in range(groups.size): neighbor_id = groups[i] group_start = offsets[i] group_end = offsets[i+1] if i+1 < groups.size else index.size group_index = index[group_start:group_end] samples = data[group_index] # 性能测试结果 %timeit my_solution() >>> 222 ms ± 8.19 ms per loop (mean ± std. dev. of 7 runs, 1 loop each) %timeit jeromes_solution() >>> 182 ms ± 8.64 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
测试发现:当数据维度更大时,两种方案的性能差异会变得极小,大部分时间仍消耗在samples = data[group_index]这一步。目前考虑是否应该先对数组排序以匹配分组,以此进一步优化性能?
内容的提问来源于stack exchange,提问作者Sandro Martens
相关产品推荐
相关产品推荐

