如何优雅地从含NumPy数组的字典中提取子字典?
优雅实现NumPy数组字典的批量筛选
核心解法
利用字典推导式遍历原字典的所有键值对,结合通用布尔索引一次性完成所有数组的筛选,彻底避免逐个键手动赋值的繁琐操作。
优化后的完整代码
import numpy as np # 生成测试数据 t1 = np.arange(0, 30, 0.5).reshape(20, 3) t2 = np.random.rand(20, 1) # 模拟包含大量键的字典 dtest = { 'p1': t1, 'p2': t2, 'p3': np.random.rand(20), # 形状(N,)的数组 'p4': np.random.rand(20, 3) # 另一个形状(N,3)的数组 } # 生成通用布尔筛选掩码,适配(N,1)或(N,)的筛选数组 mask = dtest['p2'].ravel() > 0.7 # 一行代码批量生成子字典 sub_dtest = {key: arr[mask] for key, arr in dtest.items()}
关键细节说明
- 通用布尔掩码:用
ravel()将筛选用的数组(比如p2的(20,1)数组)转为一维布尔数组,确保能对原字典中所有维度的数组((N,3)、(N,)、(N,1))正确索引,避免维度不匹配问题。 - 字典推导式的高效性:自动遍历原字典的所有键值对,批量生成筛选后的数组,无论字典有多少键,都只需一行代码完成。
- 简化索引逻辑:直接使用布尔数组作为索引,比
np.where返回的索引元组更直观,无需处理元组解包的额外操作,降低代码复杂度。
内容的提问来源于stack exchange,提问作者vin60
相关产品推荐
相关产品推荐

