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

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() 

优化方案

核心优化思路

  1. 预定义邻域偏移:所有点的邻域都是相对自身的固定偏移,提前生成偏移量避免重复计算,彻底移除meshgrid调用。
  2. 向量化批量计算:用numpy一次性生成所有点的邻域坐标,替代循环逐个计算,大幅提升单线程效率。
  3. 多进程并行:用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 18:54:54