如何用列表推导式优雅收集二维NumPy数组中的相同行?
优雅用列表推导式分组NumPy数组的相同行索引
你的需求是把二维NumPy数组中相同行的索引分组,之前的列表推导式之所以没得到预期结果,是因为它为每个行单独生成了匹配的索引列表,导致出现重复分组(比如行0和行3会互相出现在对方的列表里),而且每个列表只包含当前行之外的索引,不是完整的分组。
咱们可以换个思路:先将每行转换成可哈希的元组(因为NumPy数组不能直接作为字典键),用字典收集每个行对应的所有索引,最后用列表推导式提取这些分组即可。这种方法不仅简洁,效率也远高于你之前的循环实现(时间复杂度从O(n²)降到O(n)),完全适配你提到的最大1000×700的数据集。
实现代码
import numpy as np A = np.array([ [1,1,1,0,0,0], [0,0,1,0,1,1], [0,0,1,0,1,1], [1,1,1,0,0,0], [0,0,1,0,1,1], [1,0,0,0,0,0], [1,0,1,0,0,0], [1,0,0,0,1,0], [1,0,0,0,0,0] ]) # 第一步:用字典分组相同行的索引 row_groups = {} for idx, row in enumerate(A): row_key = tuple(row) # 将行转为元组作为字典键 if row_key not in row_groups: row_groups[row_key] = [] row_groups[row_key].append(idx) # 第二步:用列表推导式提取最终结果 result = [group for group in row_groups.values()] print(result) # 输出:[[0, 3], [1, 2, 4], [5, 8], [6], [7]]
更简洁的写法(结合collections.defaultdict)
如果你想让代码更紧凑,可以用collections.defaultdict来简化字典初始化:
from collections import defaultdict row_groups = defaultdict(list) for idx, row in enumerate(A): row_groups[tuple(row)].append(idx) result = [group for group in row_groups.values()]
为什么这个方法更优?
- 无重复分组:每个相同行的组只会被创建一次,不会像你之前的列表推导式那样生成重复的子列表。
- 效率更高:只需要遍历数组一次,对于1000行的数据集,处理速度会比原函数快很多(原函数每次删除数组行都会触发内存复制,耗时较高)。
- 代码更简洁:逻辑清晰,用列表推导式提取结果也完全符合你想要的优雅风格。
内容的提问来源于stack exchange,提问作者user1832524
相关产品推荐
相关产品推荐

