如何用Pandas更简洁高效地合并重复行的多列值至新列?
问题:使用Pandas合并重复行的多列值到新列
我正在寻找一种使用Pandas将重复行的多列值合并到新列中的方法。以下是我目前的实现代码:
import pandas as pd df = pd.DataFrame([ {"col1": 1, "col2": 2, "col3": 3, "col4": 4, "col5": 5}, {"col1": 1, "col2": 2, "col3": 3, "col4": 6, "col5": 7}, {"col1": 1, "col2": 2, "col3": 3, "col4": 8, "col5": 8}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 101}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 102}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 100}, ]) def f(x, y, z): return list({x, y, z}) new_df_rows = [] # !!! # 需将col1、col2、col3相同的重复行合并为一行,新列包含col3、col4、col5所有值的集合 # 以下代码可以运行但用了很多繁琐的写法 # !!! df_duplicated = df[df.duplicated(["col1", "col2", "col3"], keep=False)] df_duplicated_groupby = df_duplicated.groupby(["col1", "col2", "col3"]) group_names = df_duplicated_groupby.groups.keys() for group_name in group_names: group = df_duplicated_groupby.get_group(group_name) print(group_name) print(group) new_col = list({x for l in [f(row[0], row[1], row[2]) for row in group[['col3', "col4",'col5']].to_numpy()] for x in l}) new_df_rows.append({ "col1": group_name[0], "col2": group_name[1], "col3": group_name[2], "new_col": new_col }) new_df = pd.DataFrame(new_df_rows) print(new_df.to_string())
预期结果如下:
col1 col2 col3 new_col 0 1 2 3 [3, 4, 5, 6, 7, 8] 1 1 2 10 [10, 100, 101, 102]
请问是否有更简洁、高效的Pandas实现方法来达成该需求?
简洁高效的实现方法
可以利用Pandas的groupby结合自定义聚合函数,无需手动遍历分组,代码更简洁且性能更优:
import pandas as pd df = pd.DataFrame([ {"col1": 1, "col2": 2, "col3": 3, "col4": 4, "col5": 5}, {"col1": 1, "col2": 2, "col3": 3, "col4": 6, "col5": 7}, {"col1": 1, "col2": 2, "col3": 3, "col4": 8, "col5": 8}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 101}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 102}, {"col1": 1, "col2": 2, "col3": 10, "col4": 100, "col5": 100}, ]) # 定义聚合函数:提取分组内col3、col4、col5的所有值,去重后转为列表 def aggregate_values(group): all_values = group[['col3', 'col4', 'col5']].values.flatten() return list(set(all_values)) # 仅处理重复出现的分组(行数>1),分组后应用聚合函数 filtered_df = df.groupby(['col1', 'col2', 'col3']).filter(lambda x: len(x) > 1) new_df = filtered_df.groupby(['col1', 'col2', 'col3'], as_index=False).agg( new_col=('col3', aggregate_values) ) print(new_df.to_string())
代码说明:
groupby(...).filter(lambda x: len(x) > 1):筛选出仅重复出现的分组,和原代码中df.duplicated(keep=False)的逻辑一致。agg(new_col=('col3', aggregate_values)):通过聚合函数生成新列,这里用col3作为锚点字段,实际函数会处理指定的三列数据。values.flatten():将分组内的多列数据转为一维数组,统一处理所有值。list(set(all_values)):对所有值去重后转为列表,和原代码的去重逻辑完全匹配。
如果不需要筛选重复分组,直接对所有分组合并,去掉filtered_df步骤即可。
内容的提问来源于stack exchange,提问作者mlisthenewcool
相关产品推荐
相关产品推荐

