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

如何并行化fun1函数提升性能?已尝试@njit但效果不佳

优化方案:从查找优化到并行化

1. 为什么@njit单次调用慢?

Numba的JIT编译存在首次启动开销,第一次调用函数时会完成编译工作,这个时间会被计入首次执行时长,但后续重复调用会直接复用编译好的机器码,速度会显著提升。你的测试代码仅调用了一次,所以看起来比原生代码慢,但百万次调用场景下JIT的优势会充分体现。不过当前代码还有更核心的优化点——线性查找效率过低。

2. 核心优化:用二分查找替代线性查找

你的b数组已经按b[:,0]完成排序,完全可以用二分查找快速定位i1,替代原来的while循环线性遍历。二分查找的时间复杂度为O(logN),而线性查找是O(N),对于5000个元素的数组,速度能提升数百倍。

优化后的@njit函数

import numpy as np
from numba import njit, vectorize, cuda
import time

@njit
def fun1_opt(n1, b_sorted):
    # 二分查找定位第一个大于n1的元素索引
    left = 0
    right = len(b_sorted) - 1
    i1 = len(b_sorted)  # 默认值,防止n1大于所有元素
    while left <= right:
        mid = (left + right) // 2
        if b_sorted[mid][0] > n1:
            i1 = mid
            right = mid - 1
        else:
            left = mid + 1
    
    # 计算n2(注意:实际使用需添加边界判断,确保i1-2、i1+2索引有效)
    n2 = (b_sorted[i1-2][1] * b_sorted[i1-2][0] +
          b_sorted[i1-1][1] * b_sorted[i1-1][0] +
          b_sorted[i1][1] * b_sorted[i1][0] +
          b_sorted[i1+1][1] / b_sorted[i1-1][0] +
          b_sorted[i1+2][1] / b_sorted[i1+2][0])
    return n2

3. 批量调用的预热与测试

要测试百万次调用的真实速度,必须先预热JIT(提前完成编译),再执行批量调用:

# 生成并预处理数据
a = np.random.random((5000,2))
a[:,0] = a[:,0] * 100
b = a[a[:, 0].argsort()]

# 预热JIT(首次调用完成编译)
fun1_opt(40.0, b)

# 生成百万个测试用例
n1_batch = np.random.uniform(0, 100, size=1_000_000)

# 百万次循环调用测试
start_time = time.time()
results = np.empty_like(n1_batch)
for i in range(len(n1_batch)):
    results[i] = fun1_opt(n1_batch[i], b)
print("百万次循环调用 Runtime = {} s".format(time.time() - start_time))

4. 进一步并行化:向量化或CUDA加速

如果百万次调用的速度仍不满足需求,可以用以下两种并行方案:

方案A:CPU向量化(@vectorize)

@vectorize(['float64(float64, float64[:,:])'], target='parallel')
def fun1_vectorized(n1, b_sorted):
    # 复用二分查找逻辑
    left = 0
    right = len(b_sorted) - 1
    i1 = len(b_sorted)
    while left <= right:
        mid = (left + right) // 2
        if b_sorted[mid][0] > n1:
            i1 = mid
            right = mid - 1
        else:
            left = mid + 1
    
    n2 = (b_sorted[i1-2][1] * b_sorted[i1-2][0] +
          b_sorted[i1-1][1] * b_sorted[i1-1][0] +
          b_sorted[i1][1] * b_sorted[i1][0] +
          b_sorted[i1+1][1] / b_sorted[i1-1][0] +
          b_sorted[i1+2][1] / b_sorted[i1+2][0])
    return n2

# 向量化调用(直接传入批量数组)
start_time = time.time()
results = fun1_vectorized(n1_batch, b)
print("向量化百万次调用 Runtime = {} s".format(time.time() - start_time))

方案B:GPU并行(CUDA)

注意:需配备NVIDIA GPU并安装CUDA版Numba

@cuda.jit
def fun1_cuda(n1_batch, b_sorted, results):
    # 获取当前线程索引
    idx = cuda.grid(1)
    if idx < len(n1_batch):
        n1 = n1_batch[idx]
        left = 0
        right = len(b_sorted) - 1
        i1 = len(b_sorted)
        while left <= right:
            mid = (left + right) // 2
            if b_sorted[mid][0] > n1:
                i1 = mid
                right = mid - 1
            else:
                left = mid + 1
        
        results[idx] = (b_sorted[i1-2][1] * b_sorted[i1-2][0] +
                        b_sorted[i1-1][1] * b_sorted[i1-1][0] +
                        b_sorted[i1][1] * b_sorted[i1][0] +
                        b_sorted[i1+1][1] / b_sorted[i1-1][0] +
                        b_sorted[i1+2][1] / b_sorted[i1+2][0])

# 数据移至GPU
n1_device = cuda.to_device(n1_batch)
b_device = cuda.to_device(b)
results_device = cuda.device_array_like(n1_batch)

# 设置线程块与网格大小
threads_per_block = 256
blocks_per_grid = (len(n1_batch) + threads_per_block - 1) // threads_per_block

# 执行GPU计算
start_time = time.time()
fun1_cuda[blocks_per_grid, threads_per_block](n1_device, b_device, results_device)
cuda.synchronize()  # 等待GPU计算完成
results = results_device.copy_to_host()
print("CUDA百万次调用 Runtime = {} s".format(time.time() - start_time))

关键注意事项

  • 边界处理:当前代码假设n1不会靠近数组首尾(确保i1-2≥0、i1+2<len(b)),实际使用时需添加边界判断,避免索引越界。
  • GPU数据传输开销:CUDA加速时,CPU与GPU间的数据传输存在开销,仅当数据量足够大(如百万级)时,计算收益才会超过传输成本。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 15:12:49