如何在NumPy嵌套数组中查找完整子数组对应的索引
解决方法
你需要将子数组维度聚合后做全匹配判断,而不是逐元素匹配,具体实现代码如下:
import numpy as np a = np.array([[[1 ,2], [3, 4]], [[5, 6], [7, 8]]]) b = np.array([[[5, 6], [7, 8]], [[1, 2], [3, 4]]]) # 展平每个子数组的维度,仅保留第一维作为子数组的索引维度 a_2d = a.reshape(a.shape[0], -1) b_2d = b.reshape(b.shape[0], -1) # 广播对比所有子数组,判断是否完全匹配 match_mask = (b_2d[:, np.newaxis] == a_2d).all(axis=2) # 提取匹配的a索引 res = np.where(match_mask)[1].tolist() print(res) # 输出:[1, 0]
原理解释
- 你之前用
np.where(b == a)得到的是逐元素匹配的坐标,没有做子数组维度的聚合判断,所以不符合需求。 - 这里通过
reshape把每个二维子数组压缩成一维向量,将原三维数组转换为二维数组,每行对应一个原二维子数组。 - 利用numpy的广播机制,将b的每个子数组和a的所有子数组做逐元素对比,再通过
all(axis=2)判断子数组的所有元素是否完全相等,得到匹配掩码。 - 最终通过
np.where提取匹配的a数组索引即可得到结果。
异常兼容处理
如果存在b中的子数组未在a中出现的场景,可以用如下方式处理,给未匹配的项返回默认值(比如-1):
res = [] for row in match_mask: match_idx = np.where(row)[0] res.append(match_idx[0] if len(match_idx) > 0 else -1)
内容的提问来源于stack exchange,提问作者Yeggers
相关产品推荐
相关产品推荐

