如何根据索引数组切片Numpy二维数组并保留指定维度?
解决Numpy按行提取符合条件元素并保持二维结构的问题
嘿,我完全懂你的困扰——你本来想按行提取aa中对应ab值为1的元素,并且保持每行的结构,但直接用布尔索引得到的却是一维数组对吧?这其实是Numpy布尔索引的特性导致的,它会把所有匹配的元素扁平化输出,不管原来的行列结构。
哦对了,看你给出的示例ab数组里,第二、三行只有1个1,但你补充说明每行固定有2个1,应该是示例的小失误~下面的方案都是基于每行确实有2个1的前提来写的哦。
方案一:利用高级索引精准定位
这种方法通过提取每行中符合条件的列索引,再结合行索引来获取元素,能直接得到规整的(5,2)二维数组:
import numpy as np # 你的原始数组 aa = np.array([[574, 550, 548, 545, 551], [547, 539, 539, 502, 528], [503, 530, 582, 567, 505], [590, 504, 510, 578, 525], [530, 548, 501, 580, 583]]) ab = np.array([[3, 0, 2, 1, 1], [3, 2, 2, 1, 3], [0, 3, 1, 2, 0], [1, 2, 3, 1, 3], [3, 0, 1, 1, 0]]) # 生成每行的索引,扩展为二维以便匹配列索引的形状 row_indices = np.arange(aa.shape[0])[:, np.newaxis] # 提取所有ab==1的列索引,因为每行固定2个1,直接重塑为(5,2) col_indices = np.where(ab == 1)[1].reshape(-1, 2) # 用高级索引提取元素 result = aa[row_indices, col_indices] print(result)
方案二:列表推导式,更直观易懂
如果你觉得高级索引有点绕,用列表推导式对每行单独处理也很简单。因为每行固定2个元素,Numpy会自动把结果转成规整的二维数组:
result = np.array([row[ab_row == 1] for row, ab_row in zip(aa, ab)]) print(result)
这段代码会遍历aa和ab的每一行,提取当前行中ab值为1的元素,最后组合成数组。如果每行元素数量一致,结果就是标准的二维数组;如果数量不一致,会得到object类型的数组,但你这里每行固定2个,完全不用担心这个问题。
为什么原来的方法得到一维数组?
你用的aa[ab==1]是布尔索引,Numpy在处理这种索引时,会把所有满足条件的元素从原数组中“拉出来”,按一维顺序排列,不管它们原来属于哪一行,所以最终得到的是一维数组——这是布尔索引的默认行为哦。
内容的提问来源于stack exchange,提问作者mdhoney
相关产品推荐
相关产品推荐

