You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 05:43:23