Pandas大数据集高效提取无重复ID的最高得分行方案
问题描述
我有一个规模极大的pandas DataFrame,包含group、id_a、id_b和score字段,已按score降序排列(最高分位于顶部),涵盖了id_a和id_b的所有可能组合。需要提取满足以下条件的行:
- 每个
id_a和id_b仅出现一次 - 所选行对应尽可能高的
score
示例说明
结果DataFrame包含3行,所有id_a和id_b均仅出现一次:
- A2/B2和A1/B1选取了各自得分最高的行
- A3的最优匹配行对应的B1已被使用,因此选取了次优的B3组合行
输入表格
| group | id_a | id_b | score |
|---|---|---|---|
| 1 | A2 | B2 | 0.99 |
| 1 | A1 | B1 | 0.98 |
| 1 | A3 | B1 | 0.97 |
| 1 | A3 | B3 | 0.96 |
| 1 | A2 | B1 | 0.93 |
| 1 | A1 | B2 | 0.5 |
| 1 | A1 | B3 | 0.41 |
| 1 | A2 | B3 | 0.4 |
| 1 | A3 | B2 | 0.2 |
期望结果
| group | id_a | id_b | score |
|---|---|---|---|
| 1 | A2 | B2 | 0.99 |
| 1 | A1 | B1 | 0.98 |
| 1 | A3 | B3 | 0.96 |
目前通过循环遍历原始df实现需求,但处理大数据集时速度极慢。尝试过的思路遇到困难:
- 选定某行后移除已使用的ID,但不确定如何在不重启循环的情况下实现
- 按
group拆分数据(示例仅含1个分组,但实际会有大量分组,且ID在分组间唯一),但似乎无法提升效率
现有实现代码:
import pandas as pd # 创建示例DataFrame group = [1, 1, 1, 1, 1, 1, 1, 1, 1] id_a = ['A2', 'A1', 'A3', 'A3', 'A2', 'A1', 'A1', 'A2', 'A3'] id_b = ['B2', 'B1', 'B1', 'B3', 'B1', 'B2', 'B3', 'B3', 'B2'] score = [0.99, 0.98, 0.97, 0.96, 0.93, 0.5, 0.41, 0.4, 0.2] df = pd.DataFrame({'group': group, 'id_a': id_a, 'id_b': id_b, 'score': score}) result = pd.DataFrame(columns=df.columns) # 提取目标行 for i, row in df.iterrows(): if len(result) == 0: result = row.to_frame().T else: if ((row['id_a'] in result['id_a'].tolist()) or (row['id_b'] in result['id_b'].tolist())): continue else: result = pd.concat([result, row.to_frame().T[result.columns]])
高效解决方案
核心优化点
原始代码慢的主要原因:
- 使用
iterrows()遍历行,本身效率较低 - 每次判断ID是否已使用时,用
list的in操作(时间复杂度O(n)),随着结果集增大,速度会越来越慢
方案1:用集合跟踪已使用ID + 高效遍历
用set存储已使用的id_a和id_b(in操作时间复杂度O(1)),同时改用itertuples()遍历行(比iterrows()快很多):
import pandas as pd # 创建示例DataFrame group = [1, 1, 1, 1, 1, 1, 1, 1, 1] id_a = ['A2', 'A1', 'A3', 'A3', 'A2', 'A1', 'A1', 'A2', 'A3'] id_b = ['B2', 'B1', 'B1', 'B3', 'B1', 'B2', 'B3', 'B3', 'B2'] score = [0.99, 0.98, 0.97, 0.96, 0.93, 0.5, 0.41, 0.4, 0.2] df = pd.DataFrame({'group': group, 'id_a': id_a, 'id_b': id_b, 'score': score}) used_a = set() used_b = set() selected_rows = [] # 遍历行,优先选择高分且未使用的ID组合 for row in df.itertuples(index=False): if row.id_a not in used_a and row.id_b not in used_b: selected_rows.append(row) used_a.add(row.id_a) used_b.add(row.id_b) # 如果所有ID都已匹配,可提前终止循环 if len(used_a) == df['id_a'].nunique() and len(used_b) == df['id_b'].nunique(): break # 转换为结果DataFrame result = pd.DataFrame(selected_rows, columns=df.columns) print(result)
方案2:按分组批量处理
由于实际场景中存在大量分组且ID在分组间唯一,可以结合groupby和上述高效遍历逻辑,实现分组并行处理:
import pandas as pd from multiprocessing import Pool def process_group(group_df): used_a = set() used_b = set() selected_rows = [] for row in group_df.itertuples(index=False): if row.id_a not in used_a and row.id_b not in used_b: selected_rows.append(row) used_a.add(row.id_a) used_b.add(row.id_b) if len(used_a) == group_df['id_a'].nunique() and len(used_b) == group_df['id_b'].nunique(): break return pd.DataFrame(selected_rows, columns=group_df.columns) # 创建多分组示例DataFrame group = [1]*9 + [2]*9 id_a = ['A2', 'A1', 'A3', 'A3', 'A2', 'A1', 'A1', 'A2', 'A3'] + ['A4', 'A5', 'A6', 'A6', 'A4', 'A5', 'A5', 'A4', 'A6'] id_b = ['B2', 'B1', 'B1', 'B3', 'B1', 'B2', 'B3', 'B3', 'B2'] + ['B5', 'B4', 'B4', 'B6', 'B4', 'B5', 'B6', 'B6', 'B5'] score = [0.99, 0.98, 0.97, 0.96, 0.93, 0.5, 0.41, 0.4, 0.2] + [0.99, 0.98, 0.97, 0.96, 0.93, 0.5, 0.41, 0.4, 0.2] df = pd.DataFrame({'group': group, 'id_a': id_a, 'id_b': id_b, 'score': score}) # 按group拆分 groups = [group_df for _, group_df in df.groupby('group')] # 多进程处理 with Pool() as pool: results = pool.map(process_group, groups) # 合并结果 final_result = pd.concat(results, ignore_index=True) print(final_result)
方案说明
- 两种方案都采用贪心策略:因为原DataFrame已按
score降序排列,优先选择高分行,确保满足条件的同时得分尽可能高 - 使用
set存储已用ID,将判断操作的时间复杂度从O(n)降到O(1),大幅提升效率 - 多进程分组处理适合超大规模数据集,充分利用多核CPU资源
内容的提问来源于stack exchange,提问作者Jaccar
相关产品推荐
相关产品推荐

