如何编写可适配任意维度numpy数组的掩码范围扩大通用函数
任意维度numpy掩码逐行膨胀实现方案
核心用numpy.apply_along_axis实现,无需判断输入维度、无跨行污染、代码简洁:
实现代码
import numpy as np def dilate_mask(mask: np.ndarray, expand_radius: int) -> np.ndarray: # 单一行的1D膨胀逻辑 def _dilate_1d(arr): kernel = np.ones(2 * expand_radius + 1, dtype=np.int32) # same模式卷积保证输出长度和输入一致 conv_res = np.convolve(arr, kernel, mode="same") return (conv_res > 0).astype(arr.dtype) # 沿最后一个轴(行维度)自动应用到所有行,适配任意维度输入 return np.apply_along_axis(_dilate_1d, axis=-1, arr=mask)
效果测试
1D数组测试
input_1d = np.array([0,0,0,0,0,1,0,0,0,0,1,0,0,0]) print(dilate_mask(input_1d, 1)) # 输出:[0 0 0 0 1 1 1 0 0 1 1 1 0 0] print(dilate_mask(input_1d, 2)) # 输出:[0 0 0 1 1 1 1 1 1 1 1 1 1 0]
2D数组测试
input_2d = np.array([ [0,0,0,1,0,0,1,0], [0,1,0,0,0,0,0,0], [0,0,0,0,0,0,1,0] ]) print(dilate_mask(input_2d, 1)) # 输出: # [[0 0 1 1 1 1 1 1] # [1 1 1 0 0 0 0 0] # [0 0 0 0 0 1 1 1]]
方案优势
- 自动适配任意维度输入:1D、2D、3D甚至更高维数组都可直接传入,无需额外判断维度、写分支逻辑
- 无跨行污染:所有膨胀操作都在单个一维行内完成,不会跨不同行产生数据干扰
- 实现简洁:核心逻辑仅数行,无递归、无复杂的索引操作
如果你需要处理的是其他轴而不是最后一维,只要修改axis参数对应的值即可。
内容的提问来源于stack exchange,提问作者Thibs
相关产品推荐
相关产品推荐

