如何大幅加速多Numpy数组的并行滑动窗口近邻索引提取?
问题描述
我编写了一个循环查找当前索引的8个最近邻索引的脚本,本质是滑动窗口算法,需在多个维度相同(如2800i×1200j)的2-D数组上并行执行。测试时使用了12个Float32类型、最大8位小数精度的数组。
近邻查找部分脚本如下:
import numpy as np import multiprocessing as mpr def get_neighbors(arr, origin, num_neighbors = 8): coords = np.array([[i,j] for (i,j),value in np.ndenumerate(arr)]).reshape(arr.shape + (2,)) distances = np.linalg.norm(coords - origin, axis = -1) neighbor_limit = np.sort(distances.ravel())[num_neighbors] window = np.where(distances <= neighbor_limit) exclude_window = np.where(distances > neighbor_limit) return window, exclude_window, distances
我构建了静态数组gridranger处理滑动窗口索引循环,为其他同尺寸数组提供窗口索引坐标。脚本目标是提取所有数组中滑动窗口索引对应的数值至列表,再进行分析,这部分代码如下:
def extractor(queue, gridin, windowin): extract_values = [] for i in range(0, len(windowin[0])): extract_values.append(gridin[windowin[0][i], windowin[1][i]]) queue.put(extract_values) def parallel(): for index, val in np.ndenumerate(gridranger): window, exclude, distances = get_neighbors(gridranger, [index[0], index[1]]) outarr = np.column_stack((window[0], window[1])) outvalues, processes = [], [] q = mpr.Queue() for grid in grids: pro = mpr.Process(target=extractor, args=(q, grid, window)) processes.extend([pro]) pro.start() for p in processes: extract_values = q.get() outvalues.append(extract_values) for p in processes: p.join() # return outvalues print(index, outvalues)
当前使用multiprocess运行耗时约7.5-8.5秒,针对这类大型2-D数组的滑动窗口处理效率极低,请问可采取哪些步骤大幅缩短运行时间?
优化方案
1. 彻底重构邻点查找逻辑,砍掉全局计算开销
当前get_neighbors每次都遍历整个数组生成所有坐标,这是最大性能瓶颈。对于常规网格来说,8个最近邻必然在当前点的3x3局部范围内,完全不需要全局计算:
- 直接生成当前点周围的候选坐标,边界点做裁剪处理
- 用平方距离替代
np.linalg.norm,避免开根号的计算开销 - 优化后的示例代码:
def get_neighbors_fast(arr_shape, origin, num_neighbors=8): i, j = origin # 限定3x3的局部范围,边界点自动裁剪 min_i, max_i = max(0, i-1), min(arr_shape[0]-1, i+1) min_j, max_j = max(0, j-1), min(arr_shape[1]-1, j+1) # 生成候选坐标网格 i_grid, j_grid = np.meshgrid(np.arange(min_i, max_i+1), np.arange(min_j, max_j+1), indexing='ij') coords = np.stack([i_grid.ravel(), j_grid.ravel()], axis=1) # 计算平方距离(避免开根号,速度提升明显) sq_distances = ((coords - origin)**2).sum(axis=1) # 排序取前num_neighbors个(排除自身) sorted_indices = np.argsort(sq_distances) # 跳过原点(自身) start_idx = 1 if sq_distances[sorted_indices[0]] == 0 else 0 selected_indices = sorted_indices[start_idx:start_idx+num_neighbors] # 转换为窗口索引格式 return (coords[selected_indices, 0], coords[selected_indices, 1]), None, None
2. 复用进程池,避免重复创建进程的巨大开销
当前代码每个索引都要创建12个新进程,进程的创建和销毁是极大的资源浪费:
- 提前创建全局进程池,复用进程资源
- 用
starmap批量提交提取任务,替代手动管理进程队列 - 优化后的并行逻辑示例:
def extractor_fast(gridin, windowin): # 用numpy直接索引替代Python循环,提取速度提升数倍 return gridin[windowin[0], windowin[1]].tolist() def parallel_optimized(gridranger, grids): arr_shape = gridranger.shape # 根据CPU核心数创建进程池,按需调整大小 with mpr.Pool(processes=mpr.cpu_count()) as pool: for index, val in np.ndenumerate(gridranger): window, _, _ = get_neighbors_fast(arr_shape, [index[0], index[1]]) # 批量提交所有数组的提取任务 outvalues = pool.starmap(extractor_fast, [(grid, window) for grid in grids]) print(index, outvalues)
3. 优化数组内存布局,提升缓存命中率
- 把所有输入数组转换为C连续内存布局:
grid = np.ascontiguousarray(grid),numpy对连续数组的索引操作效率更高 - 保持
float32类型不变,减少内存占用,提升缓存利用率
4. 砍掉不必要的计算和数据传递
- 如果后续没有用到
exclude_window和distances,直接从函数返回值中移除,减少内存分配和数据传递 - 当前代码中
outarr变量未使用,直接删除
5. 预处理全局邻点表(仅适用于规则网格)
如果是完全规则的网格,可以一次性预处理所有点的8邻索引,后续循环直接查表即可:
def precompute_all_neighbors(arr_shape): all_neighbors = np.zeros((arr_shape[0], arr_shape[1], 8, 2), dtype=np.int32) for i in range(arr_shape[0]): for j in range(arr_shape[1]): neighbors = [] # 遍历8个方向 for di in (-1,0,1): for dj in (-1,0,1): if di == 0 and dj == 0: continue ni, nj = i+di, j+dj # 边界点判断 if 0 <= ni < arr_shape[0] and 0 <= nj < arr_shape[1]: neighbors.append([ni, nj]) # 边界点不足8个的情况,用自身填充(可按需调整) while len(neighbors) < 8: neighbors.append([i,j]) all_neighbors[i,j] = np.array(neighbors) return all_neighbors
- 预处理一次后,循环时直接通过
all_neighbors[i,j]获取邻点坐标,完全避免重复计算开销
内容的提问来源于stack exchange,提问作者TornadoEric
相关产品推荐
相关产品推荐

