如何优化Python中耗时15-20分钟的数组覆盖组合生成函数?
优化思路与实现方案
你的问题核心是嵌套循环过多导致计算量爆炸,加上每次循环都创建numpy数组的额外开销,直接拖慢了整个函数。我们可以从减少计算冗余和替换低效操作两个方向入手,大幅提升性能。
先指出原代码的两个小笔误:
- 代码里
a[z1s:z1s+z] += 1应该是z1而不是z,否则z1的长度参数完全没用到; yield里的第一个z2s应该是z2,属于变量名输入错误。
接下来是具体优化方案:
核心优化:用位掩码替代numpy数组操作
长度为12的数组刚好可以用一个12位的整数表示覆盖状态——每一位对应数组的一个元素,1表示被覆盖,0表示未覆盖。位运算(按位或、按位异或)是CPU原生支持的高效操作,比numpy数组的切片累加快几个数量级。
比如,一个从start开始、长度为length的区间,对应的掩码可以用((1 << length) - 1) << start计算:
(1 << length) -1生成连续length个1的二进制数;- 左移
start位,把这组1移动到对应的数组位置。
检查是否完全覆盖只需判断掩码是否等于0b111111111111(即十进制的4095)。
优化后的代码实现
我们加入剪枝逻辑,在循环过程中提前判断当前覆盖状态是否有可能通过后续区间补全,避免无用计算:
import itertools def precompute_masks(min_len=4, max_len=9, total_length=12): """预计算所有可能的区间对应的长度、起始位置、掩码""" masks = {} for length in range(min_len, max_len + 1): max_start = total_length - length if max_start < 0: continue template = (1 << length) - 1 # 生成连续length个1的模板 mask_list = [] for start in range(max_start + 1): mask = template << start mask_list.append((length, start, mask)) masks[length] = mask_list return masks def find_all_optimized(): target_mask = (1 << 12) - 1 # 12位全1,代表完全覆盖 masks = precompute_masks() possible_lengths = list(masks.keys()) # 枚举5个区间的长度组合 for len_comb in itertools.product(possible_lengths, repeat=5): x_len, y_len, z_len, z1_len, z2_len = len_comb # 获取对应长度的所有区间掩码信息 x_options = masks[x_len] y_options = masks[y_len] z_options = masks[z_len] z1_options = masks[z1_len] z2_options = masks[z2_len] # 遍历x和y的组合,提前剪枝 for x_info, y_info in itertools.product(x_options, y_options): x_l, x_s, x_mask = x_info y_l, y_s, y_mask = y_info current_mask = x_mask | y_mask # 计算未覆盖的位数,判断后续3个区间能否补全 uncovered = target_mask ^ current_mask uncovered_count = bin(uncovered).count('1') max_possible_cover = z_len + z1_len + z2_len if uncovered_count > max_possible_cover: continue # 后续区间不可能补全,直接跳过 # 遍历z的组合,继续剪枝 for z_info in z_options: z_l, z_s, z_mask = z_info current_mask2 = current_mask | z_mask uncovered2 = target_mask ^ current_mask2 uncovered_count2 = bin(uncovered2).count('1') max_possible_cover2 = z1_len + z2_len if uncovered_count2 > max_possible_cover2: continue # 遍历z1的组合 for z1_info in z1_options: z1_l, z1_s, z1_mask = z1_info current_mask3 = current_mask2 | z1_mask uncovered3 = target_mask ^ current_mask3 uncovered_count3 = bin(uncovered3).count('1') max_possible_cover3 = z2_len if uncovered_count3 > max_possible_cover3: continue # 遍历z2的组合,检查最终覆盖情况 for z2_info in z2_options: z2_l, z2_s, z2_mask = z2_info total_mask = current_mask3 | z2_mask if total_mask == target_mask: # 生成计数数组(如果不需要可以省略,进一步提速) a = [0] * 12 for i in range(12): cnt = 0 if x_mask & (1 << i): cnt += 1 if y_mask & (1 << i): cnt += 1 if z_mask & (1 << i): cnt += 1 if z1_mask & (1 << i): cnt += 1 if z2_mask & (1 << i): cnt += 1 a[i] = cnt yield x_l, y_l, z_l, z1_l, z2_l, x_s, y_s, z_s, z1_s, z2_s, a # 测试性能 %time list(find_all_optimized())
扩展到6个区间的方案
如果需要支持最多6个区间,只需要:
- 把
itertools.product(possible_lengths, repeat=5)改成repeat=6; - 增加对应长度的变量(比如
z3_len)和掩码选项; - 调整剪枝逻辑中的后续区间最大覆盖能力计算(比如把
max_possible_cover改成z_len + z1_len + z2_len + z3_len)。
为什么这个方案更快?
- 位运算替代数组操作:位掩码的计算和判断都是纳秒级的,远快于numpy数组的内存分配、切片累加和遍历检查;
- 多层剪枝:提前过滤掉不可能补全覆盖的组合,减少了大量无用迭代;
- 预计算掩码:避免重复生成区间模板,减少冗余计算。
内容的提问来源于stack exchange,提问作者user1596433
相关产品推荐
相关产品推荐

