如何在Python中高效查找大型字典各位置最小数值对应的键
问题描述
我有如下(采样)字典A,其原始版本包含超过17000个键,每个键对应array的长度均略高于60万(所有array长度一致)。我需要为60万个位置分别找出所有array在该位置的最小数值对应的字典键。例如在下方的字典示例中,j=0时45.16672136是所有array第一个元素的最小值,因此对应返回i=3093094;同理j=1时最小数值为1.53174068,对应返回i=1157086。
A = {3093094: array([45.16672136, 1.68053313, 13.78822307, ..., 36.18798239, 36.09565274, 35.85261821]), 1156659: array([45.46286695, 1.69632425, 13.81351489, ..., 36.54544469, 36.45329774, 36.20969689]), 1156667: array([45.43970605, 1.69026244, 13.81365067, ..., 36.51934187, 36.42716964, 36.18364528]), 1156792: array([45.29956347, 1.57736575, 13.90834355, ..., 36.43079348, 36.33804273, 36.09623309]), 1157086: array([45.38149498, 1.53174068, 13.98398836, ..., 36.57985343, 36.48684657, 36.2457831 ]), 1430072: array([45.46114909, 1.58096885, 13.95459557, ..., 36.64775128, 36.55496457, 36.31324461]), 1668445: array([45.44073352, 1.5941793 , 13.92953699, ..., 36.60630965, 36.51361336, 36.27162926]), 3055958: array([45.45006118, 1.57686417, 13.95499241, ..., 36.63558996, 36.54278917, 36.30111176]), 1078241: array([45.56175847, 1.77256163, 13.75586274, ..., 36.61441986, 36.52264105, 36.27795081])}
我目前写了如下多进程实现方案,但处理耗时过长,且多进程场景下需要复制体积庞大的A会带来额外的内存开销。请问有没有符合Python规范的优雅实现方案,可以快速完成这一原本逻辑非常简单的计算需求?
import numpy as np import os from multiprocessing import Pool C = range(len(A[3093094])) def closest(All_inputs): (A,j) = All_inputs B = list(A.keys()) my_list = [A[i][j] for i in B] return(B[np.argmin(np.array(my_list))]) with Pool(processes=os.cpu_count()) as pool: results = pool.map(closest, [(A,j) for j in C])
解决方案
你原有方案的核心问题是没有利用numpy的向量化计算能力,逐列循环+多进程复制大字典的操作会带来巨量的无效开销。直接用numpy底层优化的按轴取最小值索引的操作即可完成需求,完全不需要多进程。
核心实现
import numpy as np # 提取键列表和数值二维数组,二维数组shape为(键的数量, 数组长度) keys = np.array(list(A.keys())) value_arr = np.vstack(list(A.values())) # 按列取最小值对应的行索引,直接映射到对应键 min_row_idx = np.argmin(value_arr, axis=0) results = keys[min_row_idx]
内存不足时的分块优化
如果你的设备内存不足以放下整个(17000, 600000)的大数组,可以按列分块计算,避免内存溢出:
chunk_size = 10000 # 可根据实际内存调整块大小 n_cols = value_arr.shape[1] results = [] for start in range(0, n_cols, chunk_size): end = min(start + chunk_size, n_cols) chunk_min_idx = np.argmin(value_arr[:, start:end], axis=0) results.extend(keys[chunk_min_idx].tolist())
该方案的所有计算逻辑都在numpy的C层实现,没有Python层循环开销,耗时仅为原有多进程方案的几十分之一,也不存在多进程数据复制的额外内存消耗。
内容的提问来源于stack exchange,提问作者tcokyasar
相关产品推荐
相关产品推荐

