能否按排序顺序高效遍历该问题的所有合法最优Group解决方案?
问题核心优化逻辑
首先明确约束本质:你给出的is_valid校验逻辑等价于要求选中的n个元素的part属性互不重复,即每个part最多选1个元素。你当前实现的最大性能问题是枚举了所有C(总元素数, 分组大小)种组合,其中绝大多数都不满足part不重复的约束,属于无效枚举,完全可以提前剪枝。
分步优化方案
1. 前置数据预处理
先把所有元素按part分组,每个part内部按score降序排序,这一步只需要执行一次,无需重复计算:
from collections import defaultdict # 预处理逻辑 part_to_elms = defaultdict(list) for elm in elms: part_to_elms[elm.part].append(elm) # 每个part内部按score降序排序,保证同part下排在前面的元素得分更高 for part in part_to_elms: part_to_elms[part].sort(key=lambda x: x.score[1], reverse=True) # 过滤空part,转为列表存储 parts = list(part_to_elms.values()) part_count = len(parts) if part_count < size: raise ValueError(f"只有{part_count}个不同part,无法生成大小为{size}的分组")
预处理完成后,所有我们生成的候选解天然满足part不重复的约束,不需要再调用is_valid做校验,省掉了大量无效判断。
2. 优先级队列按需生成解
你需要按得分从高到低输出解,刚好可以用最大堆(Python标准库的heapq是最小堆,存负值即可模拟最大堆)实现按需生成,不需要全量枚举所有解:
- 堆中存储候选解的状态,每次弹出得分最高的解输出
- 输出后仅生成该解对应的下一级次优候选解加入堆,用visited集合避免重复生成
- 如果只需要前K个最优解,加个计数终止即可,不需要跑完全量逻辑
3. 优化后完整代码
import heapq from collections import defaultdict def optimal_solution_iter(self, elms, size): ''' 按得分从高到低迭代输出所有合法分组 ''' # 第一步:预处理按part分组 part_to_elms = defaultdict(list) for elm in elms: part_to_elms[elm.part].append(elm) for p in part_to_elms: part_to_elms[p].sort(key=lambda x: x.score[1], reverse=True) parts = list(part_to_elms.items()) part_count = len(parts) if part_count < size: raise ValueError(f"可用part数量{part_count}不足,无法生成大小为{size}的分组") ora = self.oracle() # 计算解的优先级key,返回负值适配最小堆 def get_solution_key(selected_parts, idx_map): total_score = 0 for part_idx in selected_parts: elm_idx = idx_map[part_idx] total_score += parts[part_idx][1][elm_idx].score[1] return -total_score # 初始最优解:选得分最高的size个part的第一个元素 top_part_indices = sorted(range(part_count), key=lambda i: parts[i][1][0].score[1], reverse=True) initial_selected = tuple(sorted(top_part_indices[:size])) initial_idx_map = {p:0 for p in initial_selected} initial_key = get_solution_key(initial_selected, initial_idx_map) initial_idx_tuple = tuple(initial_idx_map[p] for p in initial_selected) heap = [] visited = set() heapq.heappush(heap, (initial_key, initial_selected, initial_idx_tuple)) visited.add((initial_selected, initial_idx_tuple)) while heap: _, selected_parts, idx_tuple = heapq.heappop(heap) # 生成当前分组直接返回,天然合法不需要校验 current_elms = [] for p, idx in zip(selected_parts, idx_tuple): current_elms.append(parts[p][1][idx]) yield Group(current_elms, ora) # 生成次优候选1:替换某个选中part的元素为同part下一个更低得分的元素 for i in range(len(selected_parts)): p = selected_parts[i] current_idx = idx_tuple[i] if current_idx + 1 < len(parts[p][1]): new_idx_list = list(idx_tuple) new_idx_list[i] = current_idx + 1 new_idx_tuple = tuple(new_idx_list) state = (selected_parts, new_idx_tuple) if state not in visited: new_key = get_solution_key(selected_parts, dict(zip(selected_parts, new_idx_tuple))) heapq.heappush(heap, (new_key, selected_parts, new_idx_tuple)) visited.add(state) # 生成次优候选2:替换某个选中的part为未选中的part,取新part的最高得分元素 unselected_parts = [p for p in range(part_count) if p not in selected_parts] for i in range(len(selected_parts)): old_p = selected_parts[i] for new_p in unselected_parts: new_selected = list(selected_parts) new_selected[i] = new_p new_selected_sorted = tuple(sorted(new_selected)) new_idx_tuple = tuple(0 for _ in new_selected_sorted) state = (new_selected_sorted, new_idx_tuple) if state not in visited: new_key = get_solution_key(new_selected_sorted, dict(zip(new_selected_sorted, new_idx_tuple))) heapq.heappush(heap, (new_key, new_selected_sorted, new_idx_tuple)) visited.add(state)
性能提升效果
- 原实现时间复杂度为
O(C(M,n)),M为总元素数,元素数量稍大就会爆炸 - 优化后实现按需生成解,如果只需要前K个最优解,时间复杂度为
O(K log K),即使全量生成所有解,也只会遍历所有合法解,不会产生任何无效枚举,性能提升至少一个数量级
内容的提问来源于stack exchange,提问作者Joshua Snider
相关产品推荐
相关产品推荐

