如何用Numpy数组掩码查找子数组在原数组中的索引?
解决Numpy中查找子数组在原数组中对应索引的问题
首先明确数组定义(假设你用的变量名是arr和sub_arr):
import numpy as np # 原数组 arr = np.array([[1, 2], [3, 4], [5, 6], [6, 5]]) # 子数组 sub_arr = np.array([[1, 2], [3, 4]])
你的代码问题分析
你写的代码存在两个核心问题:
- 语法错误:连续写了两次
np.where,属于无效语法; - 逻辑错误:使用
~取反操作,导致你实际在查找不匹配子数组的行索引,而非目标匹配索引;另外any(0)的逻辑虽能判断原数组行是否被子数组匹配,但取反后完全偏离需求。
解决方案
根据你的预期结果index=[0,1],分两种场景给出方案:
场景1:查找子数组所有行在原数组中的对应索引(非连续匹配也适用)
这种方法会找出原数组中所有与子数组任意一行完全匹配的行索引:
# 检查子数组每行与原数组每行是否完全匹配,得到形状为(子数组行数, 原数组行数)的布尔数组 match_matrix = np.all(arr == sub_arr[:, None], axis=-1) # 提取原数组中被子数组行匹配到的索引 indices = np.where(match_matrix.any(axis=0))[0] print(indices) # 输出: [0 1]
或者用更直观的列表推导方式(适合小数组):
arr_rows = [tuple(row) for row in arr] sub_rows = set(tuple(row) for row in sub_arr) indices = [i for i, row in enumerate(arr_rows) if row in sub_rows] print(indices) # 输出: [0, 1]
场景2:查找子数组作为连续片段在原数组中的起始索引(适用于子数组是原数组连续一段的情况)
如果你的子数组是原数组中连续的一段,需要找到这段的起始位置并生成连续索引:
from numpy.lib.stride_tricks import sliding_window_view # 生成原数组的滑动窗口,窗口大小等于子数组的形状 windows = sliding_window_view(arr, window_shape=sub_arr.shape) # 压缩多余维度,使窗口形状与子数组一致 windows = windows.squeeze(axis=1) # 检查每个窗口是否与子数组完全匹配 match_flags = np.all(windows == sub_arr, axis=(1, 2)) # 获取起始索引 start_idx = np.where(match_flags)[0][0] # 生成连续索引 indices = np.arange(start_idx, start_idx + sub_arr.shape[0]) print(indices) # 输出: [0 1]
内容的提问来源于stack exchange,提问作者Dinc Kirikci
相关产品推荐
相关产品推荐

