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

如何优化点筛选循环并启用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的核心原因:

  1. 使用Python动态列表的append()方法,Numba的nopython模式不支持动态类型的Python容器
  2. 依赖全局变量band_height,nopython模式禁止访问全局变量
  3. 列表推导生成可变长度数组,编译时无法确定类型与内存布局

解决方案:适配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]

关键改动说明

  1. 预分配内存:分两次遍历,第一次统计每个点的结果数量,第二次填充预分配的数组,完全避免动态列表操作
  2. 移除全局变量:将所有参数作为函数传入,符合nopython模式的要求
  3. 计算优化:旋转坐标时直接计算偏移量(无需加回原点坐标),圆锥条件用平方比较替代开根号,减少计算开销
  4. 结果结构:返回三个数组,方便后续快速查询任意点的筛选结果

修改后的代码可完全启用nopython=True模式,性能可进一步提升至1-2秒左右。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 11:05:54