Python中高效查找三维网格邻近点的优化需求
三维网格邻点获取优化方案
问题背景
我有一个均匀采样的三维网格,需要为网格上每个点获取直接邻近点(即三维空间中上下、左右、前后6个方向的相邻点)。现有代码处理大尺寸网格(如201×201×24)时耗时极长,经cProfile分析,get_neighbour_indices方法中的meshgrid是主要性能瓶颈;尝试多线程优化因GIL限制无法实现真正并行,需寻求高效且支持并行的解决方案。
现有代码(原实现):
import numpy as np import cProfile class Neighbours: # Get neighbors @classmethod def get_neighbour_indices(cls, row, col, frame, distance=1): # Define the indices for the neighbor pixels r = np.linspace(row - distance, row + distance, 2 * distance + 1) c = np.linspace(col - distance, col + distance, 2 * distance + 1) f = np.linspace(frame - distance, frame + distance, 2 * distance + 1) nc, nr, nf = np.meshgrid(c, r, f) neighbors = np.vstack((nr.flatten(), nc.flatten(), nf.flatten())).T # Filter out valid neighbor indices within the array bounds valid_indices = (neighbors[:, 0] >= 0) & (neighbors[:, 0] < nRows) & (neighbors[:, 1] >= 0) & (neighbors[:, 1] < nCols) & (neighbors[:, 2] >= 0) & (neighbors[:, 2] < nFrames) # Return the valid neighbor indices valid_neighbors = neighbors[valid_indices] return valid_neighbors @classmethod def MapIndexVsNeighbours(cls): neighbours_info = np.empty((nRows * nCols * nFrames), dtype=object) for frame in range(nFrames): for row in range(nRows): for col in range(nCols): neighbour_indices = cls.get_neighbour_indices(row, col, frame, distance=1) flat_idx = frame * (nRows * nCols) + (row * nCols + col) neighbours_info[flat_idx] = neighbour_indices return neighbours_info ########################------------------main()-------################## ####--run if __name__ == "__main__": nRows = 151 nCols = 151 nFrames = 24 cProfile.run('Neighbours.MapIndexVsNeighbours()', sort='cumulative') print()
优化方案
核心优化思路
- 预定义邻域偏移:所有点的邻域都是相对自身的固定偏移,提前生成偏移量避免重复计算,彻底移除
meshgrid调用。 - 向量化批量计算:用numpy一次性生成所有点的邻域坐标,替代循环逐个计算,大幅提升单线程效率。
- 多进程并行:用
multiprocessing拆分任务,绕过GIL限制,充分利用多核CPU处理大网格。
优化后代码实现
import numpy as np import cProfile from multiprocessing import Pool, cpu_count class OptimizedNeighbours: # 预定义6邻域偏移(若需27邻域可调整生成逻辑) OFFSETS = np.array([[1,0,0], [-1,0,0], [0,1,0], [0,-1,0], [0,0,1], [0,0,-1]], dtype=np.int32) @classmethod def get_all_neighbours(cls, nRows, nCols, nFrames): # 生成所有网格点的坐标数组 (N, 3),N = nRows*nCols*nFrames rows = np.repeat(np.arange(nRows), nCols * nFrames) cols = np.tile(np.repeat(np.arange(nCols), nFrames), nRows) frames = np.tile(np.arange(nFrames), nRows * nCols) all_points = np.vstack([rows, cols, frames]).T # 批量计算所有点的邻域坐标:(N, 6, 3) all_neighbours = all_points[:, None, :] + cls.OFFSETS[None, :, :] # 批量筛选有效邻点(坐标在网格边界内) valid_mask = ( (all_neighbours[:, :, 0] >= 0) & (all_neighbours[:, :, 0] < nRows) & (all_neighbours[:, :, 1] >= 0) & (all_neighbours[:, :, 1] < nCols) & (all_neighbours[:, :, 2] >= 0) & (all_neighbours[:, :, 2] < nFrames) ) # 为每个点收集有效邻点 neighbours_info = np.empty(all_points.shape[0], dtype=object) for i in range(all_points.shape[0]): neighbours_info[i] = all_neighbours[i, valid_mask[i]] return neighbours_info @classmethod def _process_chunk(cls, chunk_data): # 进程内子任务:处理部分点的邻点计算 points, nRows, nCols, nFrames = chunk_data chunk_neighbours = points[:, None, :] + cls.OFFSETS[None, :, :] valid_mask = ( (chunk_neighbours[:, :, 0] >= 0) & (chunk_neighbours[:, :, 0] < nRows) & (chunk_neighbours[:, :, 1] >= 0) & (chunk_neighbours[:, :, 1] < nCols) & (chunk_neighbours[:, :, 2] >= 0) & (chunk_neighbours[:, :, 2] < nFrames) ) return [chunk_neighbours[i, valid_mask[i]] for i in range(points.shape[0])] @classmethod def get_all_neighbours_parallel(cls, nRows, nCols, nFrames): # 生成所有网格点坐标 rows = np.repeat(np.arange(nRows), nCols * nFrames) cols = np.tile(np.repeat(np.arange(nCols), nFrames), nRows) frames = np.tile(np.arange(nFrames), nRows * nCols) all_points = np.vstack([rows, cols, frames]).T # 按CPU核心数拆分数据块 num_workers = cpu_count() chunks = np.array_split(all_points, num_workers) chunk_args = [(chunk, nRows, nCols, nFrames) for chunk in chunks] # 多进程并行处理 with Pool(num_workers) as pool: results = pool.map(cls._process_chunk, chunk_args) # 合并结果 neighbours_info = np.empty(all_points.shape[0], dtype=object) idx = 0 for res in results: neighbours_info[idx:idx+len(res)] = res idx += len(res) return neighbours_info if __name__ == "__main__": nRows = 201 nCols = 201 nFrames = 24 print("=== 测试向量化版本 ===") cProfile.run('OptimizedNeighbours.get_all_neighbours(nRows, nCols, nFrames)', sort='cumulative') print("\n=== 测试多进程版本 ===") cProfile.run('OptimizedNeighbours.get_all_neighbours_parallel(nRows, nCols, nFrames)', sort='cumulative')
扩展说明
- 27邻域支持:若需要包含斜向的全3×3×3邻域,替换
OFFSETS生成逻辑:# 生成27邻域偏移(移除自身点(0,0,0)) offsets = np.array(np.meshgrid([-1,0,1], [-1,0,1], [-1,0,1])).T.reshape(-1,3) OptimizedNeighbours.OFFSETS = offsets[np.any(offsets != 0, axis=1)] - 性能选择:小网格用向量化版本更高效;超大网格(如500×500×100)优先使用多进程版本,抵消进程通信开销后可获得接近核数的加速比。
内容的提问来源于stack exchange,提问作者skm
相关产品推荐
相关产品推荐

