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

能否按排序顺序高效遍历该问题的所有合法最优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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 06:21:01