从Numpy索引元组中提取重复索引对至新数组
处理Numpy索引元组的重复项拆分
假设我们有一个由两个Numpy数组组成的元组,每个位置的元素配对形成索引对,需要把重复的索引对从原数组中移出,存入新数组。以下是具体实现方法:
实现步骤
- 把元组中的两个数组合并为二维数组,每一行对应一个索引对
- 利用
np.unique识别唯一索引对,同时追踪每个索引对的出现情况 - 提取首次出现的所有索引对作为处理后的原数组
- 收集重复出现的索引对作为新数组
- 最后将二维数组重新拆分为元组格式
代码示例
import numpy as np # 示例输入 indices = (np.array([0, 1, 0, 1]), np.array([1, 0, 0, 0])) # 1. 合并为二维索引对数组 pairs = np.vstack(indices).T # 2. 获取唯一索引对、逆索引和出现次数 unique_pairs, inverse, counts = np.unique(pairs, axis=0, return_inverse=True, return_counts=True) # 3. 构建原数组:保留每个索引对的首次出现 mask_first = np.zeros(len(inverse), dtype=bool) for idx in range(len(unique_pairs)): # 找到当前唯一索引对的第一个出现位置 first_pos = np.where(inverse == idx)[0][0] mask_first[first_pos] = True original_pairs = pairs[mask_first] indices_processed = (original_pairs[:, 0], original_pairs[:, 1]) # 4. 构建新数组:收集所有重复出现的索引对 mask_duplicates = ~mask_first duplicate_pairs = pairs[mask_duplicates] indices_2 = (duplicate_pairs[:, 0], duplicate_pairs[:, 1]) # 输出结果 print("indices =", indices_processed) print("indices_2 =", indices_2)
运行上述代码后,输出将与期望一致:
indices = (array([0, 1, 0]), array([1, 0, 0])) indices_2 = (array([1]), array([0]))
补充说明
- 如果某个索引对出现多次,上述代码会在原数组保留1次,新数组收集剩下的所有重复项
np.unique的axis=0参数是核心,确保按完整索引对(而非单个元素)去重
内容的提问来源于stack exchange,提问作者Luluio
相关产品推荐
相关产品推荐

