如何使用numba.cuda计算两个数组的最大相对误差?
基于Numba CUDA实现全局最大相对误差计算的方案
CUDA环境下求全局最大值通常采用两级归约的方案实现:
- 第一级:每个线程块先遍历自己负责的元素,计算所有r值的块内最大值,存到一个全局的块最大值数组中
- 第二级:对块最大值数组求最大值,常规场景下直接把这个小数组拷回CPU计算即可,复杂度极低
完整实现代码
import numpy as np from numba import cuda import math # 定义每个线程块的线程数,可根据显卡配置调整为128/256/512,256为通用最优值 TPB = 256 @cuda.jit def max_relative_error_kernel(a1, a2, block_max_array): # 申请块内共享内存,用于存储当前块各线程的局部最大值 shared_r = cuda.shared.array(shape=TPB, dtype=np.float64) # 全局线程ID、块内线程ID global_tid = cuda.grid(1) local_tid = cuda.threadIdx.x # 初始化当前线程的局部最大值为负无穷 thread_max = -np.inf # 网格步长遍历所有待计算元素,更新当前线程的局部最大值 for idx in range(global_tid, a1.size, cuda.gridsize(1)): # 若a1存在0值,可在此处添加除0保护逻辑 r = abs(1 - a2[idx] / a1[idx]) if r > thread_max: thread_max = r # 将当前线程的局部最大值写入共享内存 shared_r[local_tid] = thread_max # 块内同步,确保所有线程都完成共享内存写入 cuda.syncthreads() # 块内归约计算当前块的最大值 s = TPB // 2 while s > 0: if local_tid < s: if shared_r[local_tid] < shared_r[local_tid + s]: shared_r[local_tid] = shared_r[local_tid + s] cuda.syncthreads() s = s // 2 # 块内0号线程将块最大值写入全局的块最大值数组 if local_tid == 0: block_max_array[cuda.blockIdx.x] = shared_r[0] # ------------------------------ # 调用示例 # ------------------------------ # 生成测试数组 a1 = np.random.rand(10_000_000) a2 = np.random.rand(10_000_000) # 1. 将数组拷贝到CUDA设备 ca1 = cuda.to_device(a1) ca2 = cuda.to_device(a2) # 2. 计算网格大小,块上限设为1024适配大多数显卡 blocks_per_grid = min(math.ceil(a1.size / TPB), 1024) # 申请设备内存存储每个块的最大值 d_block_max = cuda.device_array(blocks_per_grid, dtype=np.float64) # 3. 调用核函数计算 max_relative_error_kernel[blocks_per_grid, TPB](ca1, ca2, d_block_max) # 4. 块最大值数组拷回CPU,求全局最大相对误差 h_block_max = d_block_max.copy_to_host() global_max_error = h_block_max.max() # 验证结果和NumPy原生实现是否一致 numpy_result = np.abs(1 - a2 / a1).max() print(f"CUDA计算结果:{global_max_error}") print(f"NumPy原生结果:{numpy_result}") print(f"结果是否匹配:{np.allclose(global_max_error, numpy_result)}")
注意事项
- 若
a1数组存在0值,需要在核函数计算r的位置添加判断逻辑,避免除0报错 - 如果对性能要求极高,可将第二级的CPU求最大值改为再启动一个单块核函数对
d_block_max做归约,适合块数量极大的场景,普通场景下CPU计算的性能损失可以忽略 - 可根据业务精度需求将
float64改为float32,能进一步提升计算速度
内容的提问来源于stack exchange,提问作者Qiang Zhang
相关产品推荐
相关产品推荐

