将二维NumPy数组转换为标签映射行索引字典的高效方法
高效转换二维NumPy数组为标签-行索引映射字典的方法
我平时处理这类标签映射问题时,会优先用NumPy的原生矢量操作来实现——毕竟纯Python循环在面对大规模数组时效率太低了。下面给你两种实用的高效方案,适配不同场景:
方案一:批量预处理+集合去重(适合超大数组)
这种方法先通过NumPy的广播操作一次性生成所有标签对应的行索引,再用集合自动去重,整体效率很高:
import numpy as np from collections import defaultdict # 示例数组 arr = np.array([ [1, 2, 3], [1, 0, 0], [1, 3, 0] ]) # 生成重复的行索引(每个行索引重复列数次数,和展平后的标签一一对应) row_indices = np.repeat(np.arange(arr.shape[0]), arr.shape[1]) # 展平数组得到所有标签 flattened_labels = arr.ravel() # 用defaultdict收集每个标签对应的行索引(集合自动去重同一行的重复标签) label_map = defaultdict(set) for label, idx in zip(flattened_labels, row_indices): label_map[label].add(idx) # 转成有序列表(如果需要排序后的行索引) label_map = {k: sorted(v) for k in label_map}
运行后得到的结果完全符合你的需求:
print(label_map) # 输出:{0: [1, 2], 1: [0, 1, 2], 2: [0], 3: [0, 2]}
方案二:按唯一标签批量查询(代码更简洁,适合中小规模数组)
如果数组规模不大,用这种更简洁的写法也很高效,直接针对每个唯一标签用np.where定位行索引:
import numpy as np arr = np.array([ [1, 2, 3], [1, 0, 0], [1, 3, 0] ]) # 获取所有唯一标签 unique_labels = np.unique(arr) label_map = {} for label in unique_labels: # 找到所有包含该标签的行索引,用np.unique去重 rows = np.unique(np.where(arr == label)[0]) label_map[label] = rows.tolist()
额外需求:按映射长度排序
如果需要按每个标签对应的行索引数量排序,直接用sorted配合自定义key即可:
# 按行索引数量从多到少排序 sorted_map = sorted(label_map.items(), key=lambda x: len(x[1]), reverse=True) print(sorted_map) # 输出:[(1, [0, 1, 2]), (0, [1, 2]), (3, [0, 2]), (2, [0])]
效率说明
两种方案都避免了逐行遍历的低效操作:
- 方案一的预处理都是NumPy的C级操作,循环部分只是简单的键值对收集,超大数组(比如10万行以上)表现更优
- 方案二代码更易读,中小规模数组下和方案一效率差距不大
内容的提问来源于stack exchange,提问作者osm
相关产品推荐
相关产品推荐

