重构计算Jaccard Index的Python函数以提升性能
优化Jaccard指数计算函数的性能问题
我编写了一个用于计算特定类别与术语Jaccard指数的函数,公式为 J(T, C) = |DOC(T) ∩ DOC(C)| / |DOC(T) ∪ DOC(C)|。当前函数运行耗时约5秒,试过生成器、集合与列表转换等优化手段,但性能提升不明显,求重构提速方案。
相关代码如下:
def calculate_j(tid, cid) -> float: doct = dids_via_tid(tid) docc = dids_via_category(cid) it1, it2 = tee(doct) it3, it4 = tee(docc) doc_either = dids_via_either(it1, it3) doc_both = dids_via_both(it2, it4) j = float(float(len(doc_both)) / float(len(doc_either))) return j def dids_via_tid(tid) -> list: did_tids_file.seek(0) for line in did_tids_file: line_words = line.split() if any(word.startswith(tid) for word in line_words): yield line_words[0] def tids_via_dids(dids): did_tids_file.seek(0) dids_set = set(dids) for line in did_tids_file: line_words = line.split() if line_words[0] in dids_set: yield from (word.split(":")[0] for word in line_words[1:]) def dids_via_both(doct, docc) -> set: set1 = {value for value in doct} set2 = {value for value in docc} if(len(set1) > len(set2)): return set1.intersection(set2) else: return set2.intersection(set1) def dids_via_either(doct, docc) -> set: set1 = set(doct) set2 = set(docc) return set1.union(set2)
核心性能瓶颈分析
当前代码的主要问题在于:
- 重复磁盘IO:
dids_via_tid、dids_via_category(推测逻辑类似)以及后续集合转换都要反复读取文件,磁盘遍历是最大性能开销。 - 生成器的无效使用:用
tee复制生成器后再转集合,本质还是要加载所有元素到内存,既没利用生成器惰性优势,又多了一次遍历开销。
具体优化方案
1. 预加载文件数据到内存(一次性读取)
把文件内容预先解析成内存字典,后续查询直接从字典取值,彻底避免重复磁盘IO。示例如下:
from collections import defaultdict # 程序启动时执行一次,预加载所有数据 did_to_tids = {} tid_to_dids = defaultdict(set) cid_to_dids = defaultdict(set) # 对应你dids_via_category的逻辑,按需构建 # 加载文档-术语映射 with open("your_tid_file_path", "r") as f: for line in f: line_words = line.strip().split() did = line_words[0] tids = {word.split(":")[0] for word in line_words[1:]} did_to_tids[did] = tids for tid in tids: tid_to_dids[tid].add(did) # 加载文档-类别映射(根据你的dids_via_category逻辑实现) with open("your_category_file_path", "r") as f: for line in f: # 假设每行格式是 "did category_id" did, cid = line.strip().split() cid_to_dids[cid].add(did)
2. 重构计算函数,直接用内存集合操作
def calculate_j(tid, cid) -> float: doct = tid_to_dids.get(tid, set()) docc = cid_to_dids.get(cid, set()) # 用公式计算并集大小,避免重复遍历元素 intersection_count = len(doct & docc) union_count = len(doct) + len(docc) - intersection_count return intersection_count / union_count if union_count != 0 else 0.0
3. 细节优化
- 去掉不必要的类型转换:原代码中
float(float(len(doc_both)) / float(len(doc_either)))完全冗余,Python3中整数用/运算直接返回float。 - 如果文件过大无法全量加载内存,可以考虑用SQLite构建索引表,或者将倒排索引拆分成多个小文件,用二分查找快速定位目标数据,减少遍历范围。
内容的提问来源于stack exchange,提问作者gybonel
相关产品推荐
相关产品推荐

