Numpy如何筛选内部第0位元素全部相同的二级子数组?
解决方案
直接用Numpy广播机制做全向量化判断,无需任何Python循环,性能最优:
import numpy as np # 1. 提取所有二级子数组的第0位元素,得到形状为 (n, 3) 的数组 first_col = arr[:, :, 0] # 2. 生成筛选掩码:判断每个二级子数组内的3个第0位元素是否全部相等 mask = (first_col == first_col[:, [0]]).all(axis=1) # 3. 用掩码索引原数组得到结果 new_arr = arr[mask]
原理说明
first_col[:, [0]]会将原形状为(n, 3)的首元素列调整为(n, 1),利用Numpy的广播特性,可以直接和(n, 3)的first_col逐元素对比all(axis=1)要求每个二级子数组对应的对比结果全部为True,也就是该子数组的3个第0位元素完全相同- 最终得到的
mask是长度为n的布尔数组,对应原数组第一维每个元素是否符合筛选条件,直接索引即可得到结果
验证结果
按你给出的示例数组,最终得到的mask为[False True True False False],索引后得到的new_arr和你预期的输出完全一致。
内容的提问来源于stack exchange,提问作者Jivan
相关产品推荐
相关产品推荐

