如何优化点筛选循环并启用Numba的nopython模式?
三维坐标点筛选的Numba优化方案
问题背景
现有一段三维坐标点筛选逻辑,初始版本单次运行耗时约27秒;通过Numba重构并加入提前终止的反向判断后,耗时降至4-5秒,但因使用append()和可变长度列表,无法启用nopython=True模式,无法完全发挥Numba的性能潜力。
初始版本代码(耗时~27s)
import numpy as np import pandas as pd # 生成测试数据 df = pd.DataFrame({ "x" : np.random.uniform(100000.5,400000.5,69000), "y" : np.random.uniform(100000.5,300000.5,69000), "z" : np.random.uniform(-50.5,300.5, 69000), }) band_width, band_height, azi, max_distance, ang_tol = 100, 10, 60, 1000, 22.5 # 提取数组与预计算参数 x, y, z, idx = df['x'].values, df['y'].values, df['z'].values, df.index.values azi_rad = np.radians(azi) cos_azi_rad, sin_azi_rad = np.cos(azi_rad), np.sin(azi_rad) tan_ang_tol_rad = np.tan(np.radians(ang_tol)) # 初始化结果数组 res_select = np.empty_like(idx, dtype=object) for id in idx : xi, yi, zi = x[id], y[id], z[id] # 先筛选xy范围内的点,减少计算量 cond_xy_cut = ((xi-max_distance) <= x) & (x <= (xi+max_distance)) & \ ((yi-max_distance) <= y) & (y <= (yi+max_distance)) x_cut, y_cut, z_cut, idx_cut = x[cond_xy_cut], y[cond_xy_cut], z[cond_xy_cut], idx[cond_xy_cut] dx, dy = x_cut-xi, y_cut-yi # 旋转坐标 x_rotate = cos_azi_rad * dx - sin_azi_rad * dy + xi y_rotate = sin_azi_rad * dx + cos_azi_rad * dy + yi dx_rot, dy_rot = x_rotate-xi, y_rotate-yi abs_dx_rot, abs_dy_rot = np.abs(dx_rot), np.abs(dy_rot) dz = z_cut-zi # 多条件筛选 cond_x = abs_dx_rot <= band_width cond_y = (abs_dy_rot <= max_distance) & (y_rotate >= yi) cond_z = np.abs(dz) <= band_height cond_cone = (tan_ang_tol_rad * abs_dy_rot) >= np.sqrt(dx_rot**2 + dz**2) cond_sum = cond_x & cond_y & cond_z & cond_cone # 排除自身点 res_select[id] = idx_cut[cond_sum & (idx_cut != id)]
现有Numba代码的问题
原Numba代码无法启用nopython=True的核心原因:
- 使用Python动态列表的
append()方法,Numba的nopython模式不支持动态类型的Python容器 - 依赖全局变量
band_height,nopython模式禁止访问全局变量 - 列表推导生成可变长度数组,编译时无法确定类型与内存布局
解决方案:适配nopython模式的Numba重构
核心思路是预分配固定内存,避免所有Python动态对象操作,同时优化计算逻辑:
修改后的完整代码
import numpy as np import pandas as pd from numba import njit # 生成测试数据 df = pd.DataFrame({ "x" : np.random.uniform(100000.5,400000.5,69000), "y" : np.random.uniform(100000.5,300000.5,69000), "z" : np.random.uniform(-50.5,300.5, 69000), }) # 参数定义 band_width, band_height, azi, max_distance, ang_tol = 100, 10, 60, 1000, 22.5 azi_rad = np.radians(azi) cos_azi_rad = np.cos(azi_rad) sin_azi_rad = np.sin(azi_rad) tan_ang_tol_rad = np.tan(np.radians(ang_tol)) # 提取数组 x = df['x'].values y = df['y'].values z = df['z'].values idx = df.index.values len_idx = len(idx) @njit(nopython=True) def numba_select_points(idx, x, y, z, max_distance, cos_azi_rad, sin_azi_rad, band_width, band_height, tan_ang_tol_rad): len_total = len(idx) # 第一步:统计每个点符合条件的结果数量,用于预分配内存 count_per_id = np.zeros(len_total, dtype=np.int64) for id in range(len_total): xi, yi, zi = x[id], y[id], z[id] count = 0 for i in range(len_total): # 提前终止不符合条件的判断 if x[i] > (xi + max_distance) or x[i] < (xi - max_distance): continue if y[i] < (yi - max_distance) or y[i] > (yi + max_distance): continue if np.abs(z[i] - zi) > band_height: continue # 计算旋转坐标(避免重复计算xi/yi) dx = x[i] - xi dy = y[i] - yi x_rotate = cos_azi_rad * dx - sin_azi_rad * dy y_rotate = sin_azi_rad * dx + cos_azi_rad * dy # 旋转后的偏移量直接用计算结果,无需加回xi/yi if np.abs(x_rotate) > band_width: continue if np.abs(y_rotate) > max_distance or y_rotate < 0: continue # 圆锥条件用平方比较,避免开根号的性能开销 lhs_sq = (tan_ang_tol_rad * np.abs(y_rotate)) ** 2 rhs_sq = x_rotate ** 2 + (z[i] - zi) ** 2 if lhs_sq < rhs_sq: continue if id == i: continue count += 1 count_per_id[id] = count # 计算每个结果的起始索引,用于后续快速查询 starts = np.zeros(len_total + 1, dtype=np.int64) for i in range(len_total): starts[i+1] = starts[i] + count_per_id[i] # 预分配结果数组 total_results = starts[-1] res_indices = np.zeros(total_results, dtype=np.int64) res_source_ids = np.zeros(total_results, dtype=np.int64) # 存储结果对应的原始点ID # 第二步:填充结果数组 current_pos = 0 for id in range(len_total): xi, yi, zi = x[id], y[id], z[id] for i in range(len_total): if x[i] > (xi + max_distance) or x[i] < (xi - max_distance): continue if y[i] < (yi - max_distance) or y[i] > (yi + max_distance): continue if np.abs(z[i] - zi) > band_height: continue dx = x[i] - xi dy = y[i] - yi x_rotate = cos_azi_rad * dx - sin_azi_rad * dy y_rotate = sin_azi_rad * dx + cos_azi_rad * dy if np.abs(x_rotate) > band_width: continue if np.abs(y_rotate) > max_distance or y_rotate < 0: continue lhs_sq = (tan_ang_tol_rad * np.abs(y_rotate)) ** 2 rhs_sq = x_rotate ** 2 + (z[i] - zi) ** 2 if lhs_sq < rhs_sq: continue if id == i: continue res_indices[current_pos] = idx[i] res_source_ids[current_pos] = idx[id] current_pos += 1 return res_source_ids, res_indices, starts # 运行函数(首次运行会编译,后续运行直接执行机器码) res_source_ids, res_indices, starts = numba_select_points( idx, x, y, z, max_distance, cos_azi_rad, sin_azi_rad, band_width, band_height, tan_ang_tol_rad ) # 示例:获取第0个点对应的筛选结果 target_id = 0 start_idx = starts[target_id] end_idx = starts[target_id + 1] target_results = res_indices[start_idx:end_idx]
关键改动说明
- 预分配内存:分两次遍历,第一次统计每个点的结果数量,第二次填充预分配的数组,完全避免动态列表操作
- 移除全局变量:将所有参数作为函数传入,符合nopython模式的要求
- 计算优化:旋转坐标时直接计算偏移量(无需加回原点坐标),圆锥条件用平方比较替代开根号,减少计算开销
- 结果结构:返回三个数组,方便后续快速查询任意点的筛选结果
修改后的代码可完全启用nopython=True模式,性能可进一步提升至1-2秒左右。
内容的提问来源于stack exchange,提问作者Spon
相关产品推荐
相关产品推荐

