如何用np.where在MxN numpy矩阵中按行返回满足条件的元素索引?
在NumPy矩阵中按行获取满足条件的元素索引
默认情况下,np.where()返回的是两个一维数组,分别对应满足条件元素的行索引集合和列索引集合——比如你给出的示例中,np.where(a == 2)会得到(array([0, 0, 1]), array([1, 2, 0])),无法直接得到按行分组的索引格式。但可以通过以下两种方式实现你的需求:
方法一:列表推导式逐行处理
直接遍历矩阵的每一行,对每行单独调用np.where()提取符合条件的元素索引,代码简洁直观:
import numpy as np a = np.array([[1, 2, 2], [2, 3, 5]]) condition = a == 2 # 逐行获取满足条件的列索引,转成列表格式 result = [np.where(row)[0].tolist() for row in condition] print(result)
输出结果:
[[1, 2], [0]]
方法二:利用np.argwhere+np.split分组(适合大矩阵)
先通过np.argwhere()获取所有满足条件的坐标对,再按行索引拆分列索引数组,避免显式循环,效率更高:
import numpy as np a = np.array([[1, 2, 2], [2, 3, 5]]) condition = a == 2 # 获取所有满足条件的坐标(行, 列) coords = np.argwhere(condition) # 统计每行满足条件的元素个数 counts = np.bincount(coords[:, 0], minlength=a.shape[0]) # 按行数拆分列索引数组 result = np.split(coords[:, 1], np.cumsum(counts)[:-1]) # 转成列表格式(可选) result = [arr.tolist() for arr in result] print(result)
输出结果与方法一一致。
内容的提问来源于stack exchange,提问作者KidSudi
相关产品推荐
相关产品推荐

