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

如何大幅加速多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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 00:52:43