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

如何并行化/加速易并行的Numba代码?

Numba代码并行化与性能优化问题

现有代码实现

以下是我用Numba加速的代码:

import numba as nb
import numpy as np

@nb.njit(cache=True)
def find_two_largest(arr):
    # 初始化第一大和第二大元素
    if arr[0] >= arr[1]:
        largest = arr[0]
        second_largest = arr[1]
    else:
        largest = arr[1]
        second_largest = arr[0]

    # 从第三个元素开始遍历数组
    for num in arr[2:]:
        if num > largest:
            second_largest = largest
            largest = num
        elif num > second_largest:
            second_largest = num
    return largest, second_largest


@nb.njit(cache=True)
def max_bar_one(arr):
    largest, second_largest = find_two_largest(arr)
    missing_maxes = np.empty_like(arr)
    for i in range(arr.shape[0]):
        if arr[i] == largest:
            if largest != second_largest:
                missing_maxes[i] = second_largest
            else:
                missing_maxes[i] = largest  # 第一大和第二大相等时的处理
        else:
            missing_maxes[i] = largest
    return missing_maxes


@nb.njit(cache=True)
def replace_max_row_wise_add_first_delete_last(d):
    """
    对除最后一行外的每一行执行max_bar_one,第一行填充全-inf
    """
    m, n = d.shape
    result = np.empty((m, n))
    result[0] = -np.inf
    for i in range(0, m - 1):
        result[i + 1, :] = max_bar_one(d[i, :])
    return result


@nb.njit(cache=True)
def main_function(d, subcusum, j):
    temp = replace_max_row_wise_add_first_delete_last(d)
    for i1 in range(temp.shape[0]):
        for i2 in range(temp.shape[1]):
            temp[i1, i2] = max(temp[i1, i2], d[i1, i2]) + subcusum[j, i2]
    return temp

数据初始化与性能测试

我用以下方式初始化测试数据:

n = 5000
A = np.random.randint(-3, 4, (n, n)).astype(float)
cusum_rows = np.cumsum(A, axis=1)
d = np.random.randint(-3, 4, (5000, 5000))

用%timeit测试性能:

%timeit main_function(d, cusum_rows, 0)
166 ms ± 1.87 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)

并行化尝试的问题

我尝试在replace_max_row_wise_add_first_delete_last中添加parallel=True并行化循环,但代码没有提速,还出现了如下提示:

Instruction hoisting:
loop #1:
Failed to hoist the following:
dependency: $value_var.73 = getitem(value=_72call__function_11, index=$parfor__index_72.90, fn=<built-in function getitem>)

因为循环中所有调用都是独立的,这个结果不符合预期。请问这段代码是否可以并行化,或者有没有其他进一步提速的方法?


优化方案与解答

1. 修复并行化失效问题

Numba的parallel=True对循环并行化有特定要求,这里的问题出在数组切片的临时对象和循环标记上:

  • 直接使用d[i, :]会生成临时切片数组,Numba无法高效并行处理这种依赖;改用d[i]传递行视图,减少临时对象开销。
  • 用nb.prange替代普通range,显式告诉Numba这是可并行的循环。

修改后的并行版函数:

@nb.njit(cache=True, parallel=True)
def replace_max_row_wise_add_first_delete_last(d):
    m, n = d.shape
    result = np.empty((m, n))
    result[0] = -np.inf
    # 用prange开启并行循环
    for i in nb.prange(0, m - 1):
        result[i + 1, :] = max_bar_one(d[i])
    return result

2. 消除嵌套循环,用向量化操作替代

main_function中的双重循环可以用NumPy向量化操作彻底替换,大幅提升效率:

@nb.njit(cache=True)
def main_function(d, subcusum, j):
    temp = replace_max_row_wise_add_first_delete_last(d)
    # 向量化逐元素取最大值 + 广播加法,替代嵌套循环
    temp = np.maximum(temp, d) + subcusum[j]
    return temp

Numba会自动优化这类向量化操作,比显式循环的执行效率高得多。

3. 精简find_two_largest逻辑

原函数的分支可以进一步简化,减少分支预测开销:

@nb.njit(cache=True)
def find_two_largest(arr):
    n = arr.shape[0]
    if n == 1:
        return arr[0], arr[0]
    # 直接用max/min初始化,减少分支
    largest = max(arr[0], arr[1])
    second_largest = min(arr[0], arr[1])
    
    for num in arr[2:]:
        if num > largest:
            second_largest = largest
            largest = num
        elif num > second_largest:
            second_largest = num
    return largest, second_largest

4. 编译与缓存优化

  • 保持所有Numba函数的cache=True,避免重复编译开销
  • 启动时可以先用小数据集调用一次所有函数,提前完成编译,避免首次运行的编译延迟

预期性能提升

经过上述优化后,并行化循环可充分利用多核CPU,结合向量化操作的优化,整体性能预计能提升3-8倍(具体取决于CPU核心数量)。


内容的提问来源于stack exchange,提问作者Simd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 18:44:53