如何优化findsim函数执行速度?已尝试NumPy优化仍效果不佳
优化方案
核心问题分析
原代码的时间复杂度为O(n²),当n=10000时,需要执行约5000万次循环,每次还要计算两个集合的交集长度,这是速度慢的根本原因。优化方向需聚焦于减少不必要的计算、替换更高效的数据结构或算法。
具体优化手段
1. 用比特位掩码替换集合,加速交集计算
如果你的"word"字符范围有限(比如ASCII可打印字符),可以将每个集合转换为整数比特位掩码:每个字符对应一个比特位,存在该字符则置1。集合交集的长度等价于两个整数按位与结果中1的个数,用int.bit_count()(Python3.10+支持,速度远快于bin(x).count('1'))计算。
示例代码:
import collections # 预定义字符到比特位的映射(以ASCII字符为例) char_to_bit = {chr(c): 1 << c for c in range(128)} # 批量转换集合为比特掩码 root_masks = [] for s in root_words: mask = 0 for c in s: mask |= char_to_bit[c] root_masks.append(mask) def findsim_fast(root_masks): pair_dict = collections.defaultdict(list) num_words = len(root_masks) for i in range(num_words): mask_i = root_masks[i] # 提前过滤:自身字符数≤3的单词,不可能和其他单词交集>3 if mask_i.bit_count() <= 3: continue for j in range(i + 1, num_words): mask_j = root_masks[j] if mask_j.bit_count() <= 3: continue # 快速计算交集字符数 if (mask_i & mask_j).bit_count() > 3: pair_dict[i].append(j) return pair_dict
2. 构建倒排索引,减少无效比对
先建立「字符→包含该字符的单词索引列表」的倒排表,对每个单词,通过倒排表快速找到所有和它有共同字符的候选单词,再统计共同字符数,避免全量O(n²)遍历。
示例代码:
import collections # 构建倒排索引 inverted_index = collections.defaultdict(list) for idx, s in enumerate(root_words): for c in s: inverted_index[c].append(idx) def findsim_inverted(root_words, inverted_index): pair_dict = collections.defaultdict(list) num_words = len(root_words) count_cache = collections.defaultdict(int) for i in range(num_words): current_set = root_words[i] if len(current_set) <= 3: continue count_cache.clear() # 累加每个候选单词与当前单词的共同字符数 for c in current_set: for idx in inverted_index[c]: if idx > i: count_cache[idx] += 1 # 筛选符合条件的索引 for j, cnt in count_cache.items(): if cnt > 3: pair_dict[i].append(j) return pair_dict
3. 提前过滤无效单词
预处理阶段直接过滤掉字符数≤3的单词,这类单词和任何其他单词的交集长度不可能超过3,直接跳过后续所有计算,减少循环基数:
# 保留字符数>3的单词,同时记录原索引 filtered_words = [(orig_idx, s) for orig_idx, s in enumerate(root_words) if len(s) > 3] # 后续基于filtered_words处理,最终结果映射回原索引即可
4. 多进程并行计算
由于每个单词的比对逻辑独立,可利用multiprocessing拆分任务到多个CPU核心执行(Python的GIL对CPU密集型任务限制大,多进程比多线程更有效)。
示例代码框架:
import multiprocessing import collections def process_chunk(args): start_i, end_i, root_masks = args local_dict = collections.defaultdict(list) for i in range(start_i, end_i): mask_i = root_masks[i] if mask_i.bit_count() <= 3: continue for j in range(i + 1, len(root_masks)): mask_j = root_masks[j] if mask_j.bit_count() <= 3: continue if (mask_i & mask_j).bit_count() > 3: local_dict[i].append(j) return local_dict def findsim_parallel(root_masks): num_workers = multiprocessing.cpu_count() num_words = len(root_masks) chunk_size = num_words // num_workers chunks = [] for i in range(num_workers): start = i * chunk_size end = start + chunk_size if i != num_workers - 1 else num_words chunks.append((start, end, root_masks)) with multiprocessing.Pool(num_workers) as pool: results = pool.map(process_chunk, chunks) # 合并所有进程的结果 pair_dict = collections.defaultdict(list) for res in results: for k, v in res.items(): pair_dict[k].extend(v) return pair_dict
效果说明
- 比特位掩码方案能将集合交集计算速度提升数倍,因为整数位操作是Python底层优化的指令。
- 倒排索引在单词字符数少、重复字符多的场景下,能大幅减少比对次数,远快于原生O(n²)逻辑。
- 并行计算适合多核机器,可将总耗时压缩至接近单核心耗时的1/核心数。
内容的提问来源于stack exchange,提问作者steppenwolf
相关产品推荐
相关产品推荐

