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

如何快速从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 12:55:37