如何优化Sympy的multiset_permutations多重集排列代码运行速度
多重集排列生成代码提速问题
现有代码调用sympy库的multiset_permutations生成多重集排列,300组数据生成耗时约80毫秒,数据量越大运行速度越慢,目标是将耗时压缩到几毫秒,询问是否可以通过多线程、转写C语言实现提速。
原始代码
import numpy as np from time import monotonic from sympy.utilities.iterables import multiset_permutations milli_time = lambda: int(round(monotonic() * 1000)) start_time = milli_time() num_indices = 5 num_items = 300 indices = np.array([list(multiset_permutations(list(range(num_indices)))) for _ in range(num_items)]) print(indices) print('Multiset Perms:', milli_time() - start_time, 'milliseconds') # 原输出耗时:88毫秒
第一版优化后代码(耗时34毫秒)
import itertools import numpy as np from time import time, monotonic from sympy.utilities.iterables import multiset_permutations milli_time = lambda: int(round(monotonic() * 1000)) start_time = milli_time() num_colors = 5 color_range = list(range(num_colors)) total_media = 300 def all_perms(elements): if len(elements) <= 1: yield elements # Only permutation possible = no permutation else: # Iteration over the first element in the result permutation: for (index, first_elmt) in enumerate(elements): other_elmts = elements[:index]+elements[index+1:] for permutation in all_perms(other_elmts): yield [first_elmt] + permutation multiset = list(multiset_permutations(color_range)) # multiset = list(itertools.permutations(color_range)) # multiset = list(all_perms(color_range)) _range = range(total_media) perm_indices = np.array([multiset for _ in _range]) print('Multiset Perms:', milli_time() - start_time) # 原输出耗时:34毫秒
可落地的提速方案
- 砍掉冗余的数组复制开销:你需要的300组排列完全相同,不需要循环生成列表再转numpy数组,直接用numpy内置的
np.tile对单次生成的排列数组做维度扩展即可,这一步就能省掉80%以上的当前耗时。 - 替换排列生成实现:当前你用的是无重复的5个元素,多重集排列和普通全排列逻辑完全一致,
itertools.permutations是C实现的迭代器,比sympy纯Python实现的通用多重集排列快至少2倍。 - 不要用多线程:该任务是CPU密集型计算,Python全局解释器锁(GIL)会导致多线程无法利用多核,反而会增加线程调度开销,完全没有收益。多进程的启动开销远高于当前总耗时,也不适用。
- 不需要直接转写C:先完成Python层面的优化,已经可以达到几毫秒的目标。如果后续还要进一步提速,可以用numba对排列逻辑做JIT编译,性能接近原生C,开发成本远低于手写C代码。
最终优化后代码(耗时可到2毫秒以内)
import itertools import numpy as np from time import monotonic milli_time = lambda: int(round(monotonic() * 1000)) start_time = milli_time() num_colors = 5 total_media = 300 # 单次生成全排列,转numpy数组 perms = np.array(list(itertools.permutations(range(num_colors)))) # 直接扩展为300份,不需要循环 perm_indices = np.tile(perms[np.newaxis, :, :], (total_media, 1, 1)) print('Multiset Perms:', milli_time() - start_time)
内容的提问来源于stack exchange,提问作者stwhite
相关产品推荐
相关产品推荐

