高效移除一维数组公共元素:适配Numba的优化方案问询
高效移除一维数组公共元素的Numba适配问题
需求背景
- 待处理主数组:一维整数数组,长度约10-15,无重复元素、未排序,需保留原有顺序
- 目标:快速移除主数组中与另一大尺寸一维数组(长度1e5-5e5,极少达7e5)的公共元素
- 性能要求:主数组会在循环中被逐步处理,性能至关重要
现有方案的局限
使用np.setdiff1d或np.in1d可实现需求,但这两个函数无法在Numba no-python模式下编译,无法满足高性能要求。
尝试norok2的2D哈希方案遇到的问题
norok2的2D数组哈希方案比常规方法快约15倍,但适配到1D场景时遇到以下问题:
- 不理解
mul_xor_hash函数的作用,以及参数init和k是否可以任意选择 - 未添加
nb.njit装饰器时,mul_xor_hash抛出类型错误:TypeError: ufunc 'bitwise_xor' not supported for the input types... - 尝试将1D数组广播为2D后,调用
mul_xor_hash(arr2[0])抛出ValueError: new type not compatible with array - 不清楚变量
delta的作用
求助目标
如果没有更优的替代方案,如何将norok2的2D哈希方案适配为1D数组的高效实现?
测试代码
import numpy as np import numba as nb n = 500000 r = 10 arr1 = np.random.permutation(n) arr2 = np.random.randint(0, n, r) # @nb.jit def setdif1d_np(a, b): return np.setdiff1d(a, b, assume_unique=True) # @nb.jit def setdif1d_in1d_np(a, b): return a[~np.in1d(a, b)]
norok2的2D方案代码
@nb.njit def mul_xor_hash(arr, init=65537, k=37): result = init for x in arr.view(np.uint64): result = (result * k) ^ x return result @nb.njit def setdiff2d_nb(arr1, arr2): # : build `delta` set using hashes delta = {mul_xor_hash(arr2[0])} for i in range(1, arr2.shape[0]): delta.add(mul_xor_hash(arr2[i])) # : compute the size of the result n = 0 for i in range(arr1.shape[0]): if mul_xor_hash(arr1[i]) not in delta: n += 1 # : build the result result = np.empty((n, arr1.shape[-1]), dtype=arr1.dtype) j = 0 for i in range(arr1.shape[0]): if mul_xor_hash(arr1[i]) not in delta: result[j] = arr1[i] j += 1 return result
内容的提问来源于stack exchange,提问作者Ali_Sh
相关产品推荐
相关产品推荐

