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

高效移除一维数组公共元素:适配Numba的优化方案问询

高效移除一维数组公共元素的Numba适配问题

需求背景

  • 待处理主数组:一维整数数组,长度约10-15,无重复元素、未排序,需保留原有顺序
  • 目标:快速移除主数组中与另一大尺寸一维数组(长度1e5-5e5,极少达7e5)的公共元素
  • 性能要求:主数组会在循环中被逐步处理,性能至关重要

现有方案的局限

使用np.setdiff1d或np.in1d可实现需求,但这两个函数无法在Numba no-python模式下编译,无法满足高性能要求。

尝试norok2的2D哈希方案遇到的问题

norok2的2D数组哈希方案比常规方法快约15倍,但适配到1D场景时遇到以下问题:

  1. 不理解mul_xor_hash函数的作用,以及参数init和k是否可以任意选择
  2. 未添加nb.njit装饰器时,mul_xor_hash抛出类型错误:TypeError: ufunc 'bitwise_xor' not supported for the input types...
  3. 尝试将1D数组广播为2D后,调用mul_xor_hash(arr2[0])抛出ValueError: new type not compatible with array
  4. 不清楚变量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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 21:20:44