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

快速生成NumPy数组有序交集布尔掩码:为何列表推导更快?

为什么纯Python列表推导比NumPy的isin/in1d更快?及优化方法

我有两个长度相同的一维NumPy数组A和B,需要生成一个布尔数组——当A对应索引的元素存在于B中时,数组值为True,且需保留原顺序以便用于索引其他数组。若无需保留顺序,通常会将数组转成集合并用交集运算符&处理,但测试发现,纯Python列表推导的速度比NumPy内置的np.isin和np.in1d快得多,想知道这一现象的原因,以及是否有进一步提升速度的方法。

测试环境代码

import numba
import numpy as np

primes = np.array([
    2,   3,   5,   7,  11,  13,  17,  19,  23,  29,  31,  37,  41,
    43,  47,  53,  59,  61,  67,  71,  73,  79,  83,  89,  97, 101,
    103, 107, 109, 113, 127, 131, 137, 139, 149, 151, 157, 163, 167,
    173, 179, 181, 191, 193, 197, 199, 211, 223, 227, 229, 233, 239,
    241, 251, 257, 263, 269, 271, 277, 281, 283, 293, 307, 311, 313,
    317, 331, 337, 347, 349, 353, 359, 367, 373, 379, 383, 389, 397,
    401, 409, 419, 421, 431, 433, 439, 443, 449, 457, 461, 463, 467,
    479, 487, 491, 499, 503, 509, 521, 523, 541, 547, 557, 563, 569,
    571, 577, 587, 593, 599, 601, 607, 613, 617, 619, 631, 641, 643,
    647, 653, 659, 661, 673, 677, 683, 691, 701, 709, 719, 727, 733,
    739, 743, 751, 757, 761, 769, 773, 787, 797, 809, 811, 821, 823,
    827, 829, 839, 853, 857, 859, 863, 877, 881, 883, 887, 907, 911,
    919, 929, 937, 941, 947, 953, 967, 971, 977, 983, 991, 997],
    dtype=np.int64)

@numba.vectorize(nopython=True, cache=True, fastmath=True, forceobj=False)
def reverse_digits(n, base):
    out = 0
    while n:
        n, rem = divmod(n, base)
        out = out * base + rem
    return out

flipped = reverse_digits(primes, 10)

def set_isin(a, b):
    return a in b

vec_isin = np.vectorize(set_isin)

primes包含1000以内的所有质数(共168个),规模适中且固定,适合测试对比。

测试结果

In [2]: %timeit np.isin(flipped, primes)
51.3 µs ± 1.55 µs per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [3]: %timeit np.in1d(flipped, primes)
46.2 µs ± 386 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [4]: %timeit setp = set(primes)
12.9 µs ± 133 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

In [5]: %timeit setp = set(primes.tolist())
6.84 µs ± 175 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

In [6]: %timeit setp = set(primes.flat)
11.5 µs ± 54.6 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

In [7]: setp = set(primes.tolist())

In [8]: %timeit [x in setp for x in flipped]
23.3 µs ± 739 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [9]: %timeit [x in setp for x in flipped.tolist()]
12.1 µs ± 76.6 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

In [10]: %timeit [x in setp for x in flipped.flat]
19.7 µs ± 249 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [11]: %timeit vec_isin(flipped, setp)
40 µs ± 317 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [12]: %timeit np.frompyfunc(lambda x: x in setp, 1, 1)(flipped)
25.7 µs ± 418 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)

In [13]: %timeit setf = set(flipped.tolist())
6.51 µs ± 44 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

In [14]: setf = set(flipped.tolist())

In [15]: %timeit np.array(sorted(setf & setp))
9.42 µs ± 78.9 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each)

其中setp = set(primes.tolist()); [x in setp for x in flipped.tolist()]的总耗时约19微秒,明显快于NumPy内置方法。


原因分析

  1. NumPy函数的通用开销:np.isin和np.in1d是通用函数,需要处理多种场景(如不同数据类型、数组广播、重复元素、缺失值等),内部会执行排序、二分查找等操作。对于小规模数据集(如168个元素),这些通用逻辑的开销远大于向量化带来的收益。
  2. 集合查找的高效性:Python集合基于哈希表实现,x in set操作的平均时间复杂度为O(1);而NumPy的isin内部是将目标数组排序后用二分查找,每个元素的查找复杂度为O(log n),加上NumPy的内部调度开销,整体效率不如纯Python的哈希查找。
  3. 小数据集下Python循环的优势:NumPy的向量化优势在大数据集(如百万级元素)才会体现,小数据集下,Python列表推导的轻量遍历+哈希查找的组合,避免了NumPy的底层开销,速度反而更快。

进一步优化方法

1. 使用Numba编译加速

利用Numba将遍历逻辑编译为机器码,避免Python循环的解释开销,同时直接操作NumPy数组,省去转列表的步骤:

from numba import njit

@njit
def numba_isin(a, b_set):
    result = np.empty(len(a), dtype=np.bool_)
    for i in range(len(a)):
        result[i] = a[i] in b_set
    return result

# 测试
setp = set(primes.tolist())
%timeit numba_isin(flipped, setp)

2. 用np.fromiter直接生成布尔数组

避免列表转NumPy数组的额外开销,直接从迭代器生成布尔数组:

setp = set(primes.tolist())
%timeit np.fromiter((x in setp for x in flipped.tolist()), dtype=np.bool_)
# 或直接遍历NumPy数组的flat迭代器
%timeit np.fromiter((x in setp for x in flipped.flat), dtype=np.bool_)

3. 大数据集切换回NumPy方法

当数组长度增长到1e5以上时,NumPy的向量化批量处理优势会超过其通用开销,此时np.isin会比纯Python方法更快,无需再用列表推导。


内容的提问来源于stack exchange,提问作者Ξένη Γήινος

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 13:44:59