如何用NumPy快速查找矩阵中连续子矩阵的左上角索引?
使用NumPy快速查找连续子矩阵的位置
核心思路
利用NumPy的滑动窗口视图(np.lib.stride_tricks.sliding_window_view)生成Y中所有与A同维度的子矩阵,通过向量化比较完成匹配,完全避免循环,大幅提升性能。
实现步骤
- 获取矩阵维度:确定A的形状
(h, w)和Y的形状(H, W),明确Y中可容纳的子矩阵范围。 - 生成滑动窗口视图:调用
sliding_window_view生成Y的所有h×w子矩阵,得到形状为(H-h+1, W-w+1, h, w)的视图(无额外内存开销,仅复用原数组内存)。 - 向量化匹配:将滑动窗口视图与A进行广播式逐元素比较,然后在子矩阵的行、列维度上判断是否所有元素都匹配。
- 提取匹配位置:用
np.argwhere找出所有匹配结果的索引,得到子矩阵左上角的坐标数组。
代码实现
import numpy as np A = np.array([[1, 0], [0, 1]]) Y = np.array([[0, 1, 0], [0, 0, 1], [1, 0, 0]]) # 获取矩阵的行列数 h, w = A.shape H, W = Y.shape # 生成所有与A同尺寸的滑动窗口 windows = np.lib.stride_tricks.sliding_window_view(Y, window_shape=(h, w)) # 检查每个窗口是否与A完全匹配 matches = (windows == A).all(axis=(2, 3)) # 提取匹配窗口的左上角索引 indices = np.argwhere(matches) print(indices)
输出结果
[[0 1]]
额外说明
- 该方法为纯向量化实现,性能远优于循环遍历,尤其适合处理大规模矩阵。
- 若Y与A的数据类型不一致,需先统一类型(例如
A = A.astype(Y.dtype)),避免比较逻辑出错。 sliding_window_view是NumPy 1.20及以上版本的功能,旧版本可使用np.lib.stride_tricks.as_strided手动实现滑动窗口,但需注意正确计算数组的stride参数以避免内存越界。
内容的提问来源于stack exchange,提问作者mrburch
相关产品推荐
相关产品推荐

