基于Combinations与Numpy的组合筛选代码优化求助
问题描述
我需要为每个组合评分以筛选出最优组合,已经完成初步实现但代码完全未优化。当RQ、NBPIS或NBSER取值较大时,代码运行耗时过长。现有实现代码如下,请问有什么方法能让代码更快得到相同结果?
import numpy as np from itertools import combinations, combinations_with_replacement # 用户设置 RQ=['A','B','C','D','E','F','G','H'] NBPIS=3 NBSER=3 # 代码部分 Combi1=np.array(list(combinations_with_replacement(RQ,NBPIS))) Combi2=combinations_with_replacement(Combi1,NBSER) Combi3=np.array([]) Compt=0 First=0 for X in Combi2: Long=0 Compt=Compt+1 Y=np.array(X) for Z in RQ: Long=Long+1 if Z not in Y: break elif Long==len(RQ): if First==0: Combi3=Y Combi3 = np.expand_dims(Combi3, axis = 0) First=1 else: Combi3=np.append(Combi3, [Y], axis = 0) # 结果输出 print(Combi3) print(Combi3.shape) print(Compt)
优化方案与代码实现
核心性能瓶颈分析
你的代码慢主要源于三点:
- 嵌套循环逐元素检查
RQ成员,时间复杂度极高; - 反复调用
np.append动态扩展数组,每次都会触发内存重新分配,数据量越大开销越夸张; - 过早将
Combi1转为numpy数组,增加了不必要的类型转换开销。
针对性优化措施
- 集合操作替代逐元素检查:将候选组合的所有元素合并为集合,直接与
RQ的集合做等值判断,一次完成所有元素的存在性验证; - 列表预收集结果:用Python列表存储符合条件的组合,最后统一转为numpy数组,避免多次
np.append的内存损耗; - 减少中间类型转换:组合生成阶段用原生Python容器操作,最后再做numpy转换,降低额外开销。
优化后的代码
import numpy as np from itertools import combinations_with_replacement # 用户设置 RQ = ['A','B','C','D','E','F','G','H'] NBPIS = 3 NBSER = 3 rq_set = set(RQ) # 预转集合,加速判断 # 生成初始组合(用原生列表,避免过早转numpy) combi1 = list(combinations_with_replacement(RQ, NBPIS)) combi3 = [] # 用列表预收集结果 compt = 0 for x in combinations_with_replacement(combi1, NBSER): compt += 1 # 合并当前组合的所有元素为集合 current_elements = set() for sub_combi in x: current_elements.update(sub_combi) # 一次判断是否包含所有RQ元素 if current_elements == rq_set: combi3.append(x) # 最后统一转为numpy数组 combi3 = np.array(combi3) # 结果输出 print(combi3) print(combi3.shape) print(compt)
超大数据量下的进阶优化(可选)
如果RQ、NBPIS或NBSER取值极大,可以用numba对核心判断逻辑做JIT编译,进一步提升速度:
import numpy as np from itertools import combinations_with_replacement from numba import njit # 用户设置 RQ = ['A','B','C','D','E','F','G','H'] NBPIS = 3 NBSER = 3 rq_tuple = tuple(RQ) # numba对tuple支持更友好 # 预生成所有初始组合的tuple形式 combi1 = [tuple(c) for c in combinations_with_replacement(RQ, NBPIS)] @njit def check_combination(x, rq): # numba加速的元素检查逻辑 seen = set() for sub in x: for elem in sub: seen.add(elem) # 验证是否包含所有目标元素 for elem in rq: if elem not in seen: return False return True combi3 = [] compt = 0 for x in combinations_with_replacement(combi1, NBSER): compt += 1 if check_combination(x, rq_tuple): combi3.append(x) combi3 = np.array(combi3) # 结果输出 print(combi3) print(combi3.shape) print(compt)
内容的提问来源于stack exchange,提问作者Marcodeco
相关产品推荐
相关产品推荐

