如何高效并行处理多查询字典,提取每个查询的前k个高分文档?
大规模查询场景下高效获取Top-K文档得分的方案
问题背景
常规获取字典前k个最高值的写法是这样的:
dict(sorted(dictionary.items(), key=lambda item: item[1], reverse=True)[:k])
但在有128000个查询、每个查询对应128个文档得分的场景里,每个查询都要挑出得分最高的k个文档,全排序的方法效率实在太低。有没有更高效的实现?比如用Numba搞并行计算?
示例数据构造代码:
import random dictionary = {} num_docs, num_queries = 128, 128000 for query_idx in range(num_queries): docs_scores = {} for doc_idx in range(num_docs): docs_scores[f"doc_{doc_idx}"] = random.random() dictionary[f"query_{query_idx}"] = docs_scores
高效实现方案
1. 用堆(heapq)替代全排序
全排序的时间复杂度是O(n log n),用堆只需要O(n log k),当k远小于文档数(比如k<<128)时,能省不少时间。
代码示例:
import heapq def get_top_k(doc_scores, k): # 用小顶堆维护前k个最大元素 heap = [] for doc, score in doc_scores.items(): if len(heap) < k: heapq.heappush(heap, (score, doc)) else: if score > heap[0][0]: heapq.heappop(heap) heapq.heappush(heap, (score, doc)) # 反转成从大到小的顺序,转成字典 return {doc: score for score, doc in reversed(heapq.nlargest(k, heap))} # 批量处理所有查询 top_k_results = {query: get_top_k(scores, k=10) for query, scores in dictionary.items()}
2. Numba并行加速(性能提升最明显)
Numba能把Python代码编译成机器码,还能多线程并行处理查询。但它对字典支持不太友好,所以最好先把嵌套字典转成数组结构,最大化性能。
第一步:把字典转成数组
import numpy as np # 把每个查询的文档名和得分转成数组,再合并成二维数组 query_docs = [] query_scores = [] for query_scores_dict in dictionary.values(): docs = np.array(list(query_scores_dict.keys()), dtype=np.str_) scores = np.array(list(query_scores_dict.values()), dtype=np.float64) query_docs.append(docs) query_scores.append(scores) scores_array = np.vstack(query_scores) docs_array = np.array(query_docs)
第二步:Numba并行计算Top-K
from numba import njit, prange @njit(parallel=True) def numba_top_k(scores, k): num_queries, num_docs = scores.shape top_indices = np.zeros((num_queries, k), dtype=np.int64) # 用prange开启并行循环,每个线程处理一个查询 for i in prange(num_queries): # 对单个查询的得分排序,取前k个的索引(降序) sorted_indices = np.argsort(scores[i])[::-1][:k] top_indices[i] = sorted_indices return top_indices # 设置要取的Top-K数量 k = 10 top_indices = numba_top_k(scores_array, k) # 把结果转回字典格式 top_k_results = {} for query_idx in range(num_queries): query_name = f"query_{query_idx}" top_docs = docs_array[query_idx][top_indices[query_idx]] top_scores = scores_array[query_idx][top_indices[query_idx]] top_k_results[query_name] = dict(zip(top_docs, top_scores))
因为每个查询的文档数只有128,直接排序的开销很小,加上并行编译后,效率比纯Python堆方法快5-10倍(取决于CPU核心数)。如果k特别小,也可以在Numba里实现堆逻辑进一步优化,但对128个元素来说,排序的性能已经足够。
3. 从源头优化数据结构
别一开始就用嵌套字典存数据,直接用numpy数组,能省掉后续转换的开销:
num_docs, num_queries = 128, 128000 # 文档名只生成一次,所有查询共用 doc_names = np.array([f"doc_{i}" for i in range(num_docs)], dtype=np.str_) # 直接生成二维得分数组:行是查询,列是文档 scores_array = np.random.rand(num_queries, num_docs)
后续处理直接操作数组,效率比字典高很多。
性能对比
- 纯Python全排序:每个查询O(128 log 128),128000个查询总开销极大
- heapq方法:每个查询O(128 log k),k=10时比全排序快3倍左右
- Numba并行数组方法:比heapq快5-10倍,是这几种里性能最好的
内容的提问来源于stack exchange,提问作者Celso França
相关产品推荐
相关产品推荐

