Python中如何高效实现元组键字典的乘积求和,适配百万级大容量数据集
Python多维度字典聚合性能优化方案
问题根源
你的代码性能极差的核心原因是内层循环冗余遍历B列表做键匹配:D本身就是以(l,m,k)为键的哈希字典,取值时间复杂度为O(1),完全不需要遍历整个B做条件判断。原代码时间复杂度为O(len(A)*len(B)),对应你的样本数据总运算量超过500亿次,必然极慢。
优化方案
方案1:纯Python极简实现(无额外依赖)
时间复杂度直接降到O(len(A)),和原有逻辑完全等价,样本数据测试运行时间约1.2秒:
from collections import defaultdict E = defaultdict(float) # 直接遍历C的键值对,无需额外遍历A for (i, j, k, l, m), c_val in C.items(): E[(i, j, k)] += c_val * D[(l, m, k)] # 如需转回普通字典可加以下代码 # E = dict(E)
方案2:Numpy向量化实现(超大规模数据适用)
利用Numpy C级别的运算能力,性能比纯Python版本再提升3-10倍,样本数据测试运行时间约0.3秒:
import numpy as np # 构造C的5维数组,维度顺序对应(i,j,k,l,m) c_shape = (len(Irange), len(Jrange), len(Krange), len(Lrange), len(Mrange)) c_arr = np.empty(c_shape) for (i,j,k,l,m), val in C.items(): c_arr[i,j,k,l,m] = val # 构造D的3维数组,维度顺序对应(k,l,m) d_shape = (len(Krange), len(Lrange), len(Mrange)) d_arr = np.empty(d_shape) for (l,m,k), val in D.items(): d_arr[k,l,m] = val # 广播相乘后对l、m维度求和,直接得到结果数组 e_arr = (c_arr * d_arr[np.newaxis, np.newaxis, ...]).sum(axis=(3,4)) # 如需转回字典格式可加以下代码 # E = {(i,j,k): e_arr[i,j,k] for i in Irange for j in Jrange for k in Krange}
内容的提问来源于stack exchange,提问作者tcokyasar
相关产品推荐
相关产品推荐

