如何统计高维NumPy数组中低维NumPy子数组的出现次数
统计NumPy数组中子数组的出现次数
嘿,这个问题我之前做图像特征匹配的时候也碰到过,用NumPy处理这类子数组匹配其实有挺简洁的方案,我来给你一步步说明~
针对你的具体例子
首先咱们把输入转换成符合Python语法的NumPy数组(原输入里的逗号得补上,不然会报错),然后利用NumPy的滑动窗口工具快速实现匹配:
import numpy as np # 定义输入数组(补全逗号,符合Python语法规范) a = np.array([[0, 0, 1, 1], [1, 1, 1, 1], [0, 0, 1, 0], [1, 1, 1, 1]]) b = np.array([[0, 0, 1], [1, 1, 1]]) # 生成a中所有和b形状完全一致的滑动窗口 # sliding_window_view是NumPy 1.20+新增的工具,安全且易用 windows = np.lib.stride_tricks.sliding_window_view(a, window_shape=b.shape) # 检查每个窗口是否与b完全相等,最后统计匹配总数 # axis=(-2, -1)表示沿着窗口的最后两个维度(也就是b的形状维度)做全量匹配 match_count = np.sum(np.all(windows == b, axis=(-2, -1))) print(match_count) # 输出结果:2,和你的预期一致
代码逻辑拆解
sliding_window_view会把原数组a转换成一个更高维的数组,其中每一个元素都是和b形状相同的滑动窗口。比如你的例子里,a是(4,4),b是(2,3),生成的windows形状会是(3,2,2,3)——前两个维度是窗口在a中的起始位置,后两个维度是窗口本身的形状。np.all(..., axis=(-2, -1))会逐个检查每个窗口的所有元素是否和b对应相等,返回一个布尔数组,最后用np.sum统计所有True的数量,就是匹配次数。
高维数组的通用解法
不管你的数组是2维、3维还是更高维度,核心思路都是一致的:
- 用
sliding_window_view生成原数组中所有与目标子数组形状一致的滑动窗口; - 沿着子数组的维度做全量相等检查;
- 统计匹配的总数。
举个高维的例子:假设a是一个(5,5,5)的3维数组,b是一个(2,2)的2维子数组,我们要统计b在a的最后两个维度中的出现次数:
import numpy as np # 随机生成高维测试数组 a = np.random.randint(0, 2, size=(5,5,5)) b = np.array([[1,0], [0,1]]) # 生成对应形状的滑动窗口 windows = np.lib.stride_tricks.sliding_window_view(a, window_shape=b.shape) # 动态获取子数组的维度索引,适配任意维度的b match_axis = tuple(range(-len(b.shape), 0)) match_count = np.sum(np.all(windows == b, axis=match_axis)) print(match_count)
注意事项
- 如果你使用的NumPy版本低于1.20,
sliding_window_view不可用,可以用np.lib.stride_tricks.as_strided手动实现滑动窗口,但要注意计算步长时避免越界问题,推荐尽量升级到较新的NumPy版本。 - 如果子数组的匹配需要考虑重叠情况,上面的方法已经自动支持了——滑动窗口本身就是包含所有可能的重叠位置的。
内容的提问来源于stack exchange,提问作者Taufik_TF
相关产品推荐
相关产品推荐

