如何加速卷积函数?Numba优化与实现改进技术问询
问题
我编写了如下卷积函数:
import numpy as np import numba as nb # Generate sample input data num_chans = 111 num_bins = 47998 num_rad = 8 num_col = 1000 rng = np.random.default_rng() wvl_sensor = rng.uniform(low=1000, high=11000, size=(num_chans, num_col)) fwhm_sensor = rng.uniform(low=0.01, high=2.0, size=num_chans) wvl_lut = rng.uniform(low=1000, high=11000, size=(num_bins)) rad_lut = rng.uniform(low=0, high=1, size=(num_rad, num_bins)) # Original convolution implementation def original_convolve(wvl_sensor, fwhm_sensor, wvl_lut, rad_lut): sigma = fwhm_sensor / (2.0 * np.sqrt(2.0 * np.log(2.0))) var = sigma ** 2 denom = (2 * np.pi * var) ** 0.5 numer = np.exp(-(wvl_lut[:, None] - wvl_sensor[None, :])**2 / (2*var)) response = numer / denom response /= response.sum(axis=0) resampled = np.dot(rad_lut, response) return resampled
numpy版本运行耗时约45秒:
# numpy version num_chans, num_col = wvl_sensor.shape num_bins = wvl_lut.shape[0] num_rad = rad_lut.shape[0] original_res = np.empty((num_col, num_rad, num_chans), dtype=np.float64) for x in range(wvl_sensor.shape[1]): original_res[x, :, :] = original_convolve(wvl_sensor[:, x], fwhm_sensor, wvl_lut, rad_lut)
我尝试用Numba对其进行加速:
@nb.jit(nopython=True) def numba_convolve(wvl_sensor, fwhm_sensor, wvl_lut, rad_lut): num_chans, num_col = wvl_sensor.shape num_bins = wvl_lut.shape[0] num_rad = rad_lut.shape[0] output = np.empty((num_col, num_rad, num_chans), dtype=np.float64) sigma = fwhm_sensor / (2.0 * np.sqrt(2.0 * np.log(2.0))) var = sigma ** 2 denom = (2 * np.pi * var) ** 0.5 for x in nb.prange(num_col): numer = np.exp(-(wvl_lut[:, None] - wvl_sensor[None, :, x])**2 / (2*var)) response = numer / denom response /= response.sum(axis=0) resampled = np.dot(rad_lut, response) output[x, :, :] = resampled return output
但仍耗时约32秒。注意,若使用@nb.jit(nopython=True, parallel=True),输出结果全为零值。
请问如何正确使用Numba?或有哪些改进卷积函数的方法?
优化方案
1. 修复Numba并行模式的零值问题
你用parallel=True时出现零值,核心原因是并行循环中对数组切片的无保护写入引发竞态条件,且广播操作在并行环境下被错误优化。以下是修正后的并行版本:
@nb.jit(nopython=True, parallel=True) def numba_convolve_fixed(wvl_sensor, fwhm_sensor, wvl_lut, rad_lut): num_chans, num_col = wvl_sensor.shape num_bins = wvl_lut.shape[0] num_rad = rad_lut.shape[0] # 预计算常量,避免循环内重复计算 sigma = fwhm_sensor / (2.0 * np.sqrt(2.0 * np.log(2.0))) var = sigma ** 2 denom = (2 * np.pi * var) ** 0.5 inv_2var = 1.0 / (2 * var) output = np.empty((num_col, num_rad, num_chans), dtype=np.float64) # 并行遍历num_col,每个迭代完全独立 for x in nb.prange(num_col): wvl_sensor_col = wvl_sensor[:, x] response = np.empty((num_bins, num_chans), dtype=np.float64) # 显式遍历通道,消除广播的隐式开销与并行冲突 for c in range(num_chans): wvl_diff = wvl_lut - wvl_sensor_col[c] response[:, c] = np.exp(-(wvl_diff * wvl_diff) * inv_2var[c]) / denom[c] # 归一化每个通道的响应 sum_response = response.sum(axis=0) response /= sum_response resampled = np.dot(rad_lut, response) output[x, :, :] = resampled return output
优化说明:
- 预计算
inv_2var等常量,减少循环内的重复运算 - 显式遍历通道,让并行逻辑更清晰,避免Numba对广播操作的编译歧义
- 每个循环迭代的数组操作完全独立,消除竞态条件
2. 进一步降低内存开销
原代码中每次循环都会分配(num_bins, num_chans)的response数组,内存分配开销占比高。可以预分配数组重复使用:
@nb.jit(nopython=True, parallel=True) def numba_convolve_optimized(wvl_sensor, fwhm_sensor, wvl_lut, rad_lut): num_chans, num_col = wvl_sensor.shape num_bins = wvl_lut.shape[0] num_rad = rad_lut.shape[0] sigma = fwhm_sensor / (2.0 * np.sqrt(2.0 * np.log(2.0))) var = sigma ** 2 denom = (2 * np.pi * var) ** 0.5 inv_2var = 1.0 / (2 * var) output = np.empty((num_col, num_rad, num_chans), dtype=np.float64) # 预分配response数组,循环内重复使用 response = np.empty((num_bins, num_chans), dtype=np.float64) for x in nb.prange(num_col): wvl_sensor_col = wvl_sensor[:, x] for c in range(num_chans): wvl_diff = wvl_lut - wvl_sensor_col[c] response[:, c] = np.exp(-(wvl_diff * wvl_diff) * inv_2var[c]) / denom[c] sum_response = response.sum(axis=0) # 处理除以零的边界情况 for c in range(num_chans): if sum_response[c] != 0: response[:, c] /= sum_response[c] resampled = np.dot(rad_lut, response) output[x, :, :] = resampled return output
优化说明:
- 预分配
response数组,避免每次迭代的内存分配与释放开销 - 显式处理除以零的情况,避免潜在的NaN或零值问题
- 用
wvl_diff * wvl_diff替代wvl_diff ** 2,Numba对乘法的编译效率更高
3. 算法层面优化:利用高斯核的局部性
高斯函数随距离增加快速衰减,差值超过3σ的区域几乎为0,可忽略计算。通过定位有效区间大幅减少运算量:
@nb.jit(nopython=True, parallel=True) def numba_convolve_local(wvl_sensor, fwhm_sensor, wvl_lut, rad_lut): num_chans, num_col = wvl_sensor.shape num_bins = wvl_lut.shape[0] num_rad = rad_lut.shape[0] sigma = fwhm_sensor / (2.0 * np.sqrt(2.0 * np.log(2.0))) var = sigma ** 2 denom = (2 * np.pi * var) ** 0.5 inv_2var = 1.0 / (2 * var) # 预排序LUT波长,方便快速查找有效区间 sorted_wvl_idx = np.argsort(wvl_lut) sorted_wvl = wvl_lut[sorted_wvl_idx] sorted_rad_lut = rad_lut[:, sorted_wvl_idx] output = np.empty((num_col, num_rad, num_chans), dtype=np.float64) for x in nb.prange(num_col): wvl_sensor_col = wvl_sensor[:, x] for c in range(num_chans): current_wvl = wvl_sensor_col[c] current_sigma = sigma[c] # 计算有效区间:[current_wvl - 3σ, current_wvl + 3σ] lower = current_wvl - 3 * current_sigma upper = current_wvl + 3 * current_sigma # 二分查找区间边界 left = np.searchsorted(sorted_wvl, lower, side='left') right = np.searchsorted(sorted_wvl, upper, side='right') if left >= right: output[x, :, c] = 0.0 continue # 仅计算有效区间内的高斯值 wvl_diff = sorted_wvl[left:right] - current_wvl gauss_vals = np.exp(-(wvl_diff * wvl_diff) * inv_2var[c]) / denom[c] sum_gauss = gauss_vals.sum() if sum_gauss == 0: output[x, :, c] = 0.0 continue gauss_vals /= sum_gauss resampled_c = np.dot(sorted_rad_lut[:, left:right], gauss_vals) output[x, :, c] = resampled_c return output
优化说明:
- 预排序
wvl_lut,用二分查找快速定位高斯核的有效区间 - 仅计算区间内的数值,对于小σ场景,计算量可减少90%以上
- 按通道直接计算赋值,避免大矩阵乘法的内存占用
4. 其他可选优化
- GPU加速:若有GPU,用CuPy替代NumPy,只需将
np替换为cp,矩阵运算速度可提升一个数量级 - MKL优化NumPy:安装
intel-numpy,无需修改代码即可提升原生NumPy的运算速度 - 批量处理:将
num_col分块,结合Numba的向量化操作进一步提升效率
内容的提问来源于stack exchange,提问作者zxdawn
相关产品推荐
相关产品推荐

