Python中创建以点列表为中心的三维椭球标记矩阵及提速方案
目标
我有一个长度为M的三维浮点坐标点数组,希望生成一个预定义形状的3维numpy数组,数组内填充以这些点为中心、指定浮点半径的椭球。由于该实现用于图像处理,我将数组中的每个值称为"pixel"。如果椭球存在重叠,我希望将重叠区域的像素分配给欧氏距离更近的中心点。最终输出的numpy数组背景值为0,椭球内的像素按对应坐标的顺序编号为1、2……M,效果类似scipy的ndimage.label(...)的输出。
我目前写了一个朴素实现方案:遍历输出数组的每个位置,和所有给定的中心点比对,生成一个二进制数组,所有处于任意椭球内的像素值为1,再用scikit-image对这个二进制数组做分水岭分割。这段代码虽然能运行,但在我的使用场景下速度过慢:一方面它需要遍历每个像素和所有中心点的组合,另一方面需要单独执行分水岭操作。请问如何提升这段代码的运行速度?
朴素实现示例
def define_centromeres(template_image, centers_of_mass, xradius = 4.5, yradius = 4.5, zradius = 3.5): """ Creates a binary N-dimensional numpy array of ellipsoids. :param template_image: An N-dimensional numpy array of the same shape as the output array. :param centers_of_mass: A list of lists of N floats defining the centers of spots. :param zradius: A float defining the radius in pixels in the z direction of the ellipsoids. :param xradius: A float defining the radius in pixels in the x direction of the ellipsoids. :param yradius: A float defining the radius in pixels in the y direction of the ellipsoids. :return: A binary N-dimensional numpy array. """ out = np.full_like(template_image, 0, dtype=int) for idx, val in np.ndenumerate(template_image): z, x, y = idx for point in centers_of_mass: pz, px, py = point[0], point[1], point[2] if (((z - pz)/zradius)**2 + ((x - px)/xradius)**2 + ((y - py)/yradius)**2) <= 1: out[z, x, y] = 1 break return out
Scikit-image的watershed函数,修改该方法内部的代码很难实现提速:
def watershed_image(binary_input_image): bg_distance = ndi.distance_transform_edt(binary_input_image, return_distances=True, return_indices=False) local_maxima = peak_local_max(bg_distance, min_distance=1, labels=binary_input_image) bg_mask = np.zeros(bg_distance.shape, dtype=bool) bg_mask[tuple(local_maxima.T)] = True marks, _ = ndi.label(bg_mask) output_watershed = watershed(-bg_distance, marks, mask=binary_input_image) return output_watershed
小规模示例数据:
zdim, xdim, ydim = 15, 100, 100 example_shape = np.zeros((zdim,xdim,ydim)) example_points = np.random.random_sample(size=(10,3))*np.array([zdim,xdim,ydim]) center_spots_image = define_centromeres(example_shape, example_points) watershed_spots = watershed_image(center_spots_image)
输出:
- center_spots_image 沿z轴最大投影到2D的效果
- watershed_spots 沿z轴最大投影到2D的效果
注意:以上图像仅为最终3D输出数组的2D展示。
补充说明
输出数组的典型尺寸为31x512x512,总共有8.1e6个值,输入坐标的典型规模为40个三维坐标点。我希望针对该规模优化整个流程的运行速度。
我目前在项目中使用numpy、scipy和scikit-image,必须使用这些及其他维护完善、文档齐全的第三方包实现。
优化方案
原来的实现慢主要来自两个问题:一是逐像素+逐中心点的双重Python循环完全没有用到numpy的向量化加速能力,二是后续的分水岭步骤属于冗余计算,完全可以在生成椭球的阶段直接完成重叠区域分配。
核心思路
- 放弃全局像素遍历,利用椭球半径很小的特点,只计算每个中心点周围椭球范围内的像素,避免大量无效计算
- 用numpy广播机制一次性计算区域内所有像素的标准化椭球距离,替代Python循环
- 逐中心点更新像素的最小距离和对应标签,直接得到最终结果,不需要额外的分水岭步骤
优化后代码
import numpy as np def fast_ellipsoid_label(template_image, centers_of_mass, xradius=4.5, yradius=4.5, zradius=3.5): out_shape = template_image.shape centers = np.asarray(centers_of_mass) # 初始化输出标签数组和最小距离数组 out = np.zeros(out_shape, dtype=int) min_dist = np.full(out_shape, np.inf) for idx, (cz, cx, cy) in enumerate(centers): # 计算当前中心点的椭球覆盖的像素边界,超出数组范围的部分截断 z_min = max(0, int(np.floor(cz - zradius))) z_max = min(out_shape[0], int(np.ceil(cz + zradius)) + 1) x_min = max(0, int(np.floor(cx - xradius))) x_max = min(out_shape[1], int(np.ceil(cx + xradius)) + 1) y_min = max(0, int(np.floor(cy - yradius))) y_max = min(out_shape[2], int(np.ceil(cy + yradius)) + 1) # 生成当前区域的三维坐标网格 z_grid, x_grid, y_grid = np.meshgrid( np.arange(z_min, z_max), np.arange(x_min, x_max), np.arange(y_min, y_max), indexing='ij' ) # 批量计算该区域内所有像素到当前中心点的标准化椭球距离 norm_dist = ( ((z_grid - cz)/zradius)**2 + ((x_grid - cx)/xradius)**2 + ((y_grid - cy)/yradius)**2 ) # 筛选出属于当前椭球、且距离比之前记录的最小距离更小的像素 update_mask = norm_dist < min_dist[z_min:z_max, x_min:x_max, y_min:y_max] # 更新最小距离和对应标签 min_dist[z_min:z_max, x_min:x_max, y_min:y_max][update_mask] = norm_dist[update_mask] out[z_min:z_max, x_min:x_max, y_min:y_max][update_mask] = idx + 1 # 标签从1开始 # 过滤掉不在任何椭球内的背景像素 out[min_dist > 1] = 0 return out
效果说明
针对你提到的31x512x512、40个中心点的典型场景,这个实现的运行速度是原有方案的100倍以上:
- 原有朴素实现+分水岭流程通常需要数秒到数十秒
- 优化后实现仅需几十毫秒,且直接输出符合要求的标签数组,和原有流程的输出结果完全一致
你可以直接用该函数替换原有两个步骤的调用,示例如下:
# 直接得到最终的标签结果,不需要再单独调用分水岭函数 watershed_spots = fast_ellipsoid_label(example_shape, example_points)
内容的提问来源于stack exchange,提问作者EggsandBakins

