如何基于列中部分重叠值合并PySpark DataFrame行?
基于PySpark DataFrame合并共享来源的行
需求:当DataFrame的行之间共享sources列中的值时,合并values集合与去重后的sources集合;无共享来源的行保持不变。已完成sources列的explode操作,寻求后续实现方法。
原DataFrame示例
| values | sources |
|---|---|
| [a, b] | [s1, s2, s3] |
| [b, c, d] | [s5, s1] |
| [x, y] | [s7] |
期望输出
| values | sources |
|---|---|
| [a, b, c, d] | [s1, s2, s3, s5] |
| [x, y] | [s7] |
已执行代码
exploded_df = df.select( col("values").alias("values"), explode(col("sources")).alias("source") ).select( col("values"), col("source") )
执行后得到的DataFrame
| values | source |
|---|---|
| [a, b] | s1 |
| [a, b] | s2 |
| [a, b] | s3 |
| [b, c, d] | s5 |
| [b, c, d] | s1 |
| [x, y] | s7 |
解决方案
这本质是连通分量问题:共享source的行属于同一个连通组,需合并组内所有数据。以下是两种实现方式:
方式一:使用GraphFrames(推荐,高效)
GraphFrames是Spark的图计算库,适合处理这类连通性问题,需先安装:pip install graphframes
from pyspark.sql import functions as F from graphframes import GraphFrame # 1. 给原始DataFrame添加唯一行ID,用于关联 df_with_id = df.withColumn("row_id", F.monotonically_increasing_id()) # 2. 构建图的顶点和边 # 顶点:每个唯一行ID作为顶点 vertices = df_with_id.select("row_id").distinct() # 边:同一source对应的不同行ID之间建立连接 edges = exploded_df.join(df_with_id, on="values", how="inner") \ .select(F.col("row_id").alias("src"), F.col("row_id").alias("dst"), "source") \ .filter(F.col("src") != F.col("dst")) # 3. 计算连通分量 g = GraphFrame(vertices, edges) connected_components = g.connectedComponents() # 4. 按连通组分组合并数据 result_df = df_with_id.join(connected_components, on="row_id", how="inner") \ .groupBy("component") \ .agg( # 合并并去重values F.array_distinct(F.flatten(F.collect_set("values"))).alias("values"), # 合并并去重sources F.array_distinct(F.flatten(F.collect_set("sources"))).alias("sources") ) \ .drop("component") result_df.show(truncate=False)
方式二:无依赖迭代实现(适合小数据)
如果无法安装GraphFrames,可通过迭代合并连通组实现,效率较低,仅推荐小数据集使用:
from pyspark.sql import functions as F from itertools import chain # 1. 给原始DataFrame添加唯一行ID df_with_id = df.withColumn("row_id", F.monotonically_increasing_id()) # 2. 收集每个source对应的行ID集合 source_row_map = exploded_df.join(df_with_id, on="values", how="inner") \ .groupBy("source") \ .agg(F.collect_set("row_id").alias("row_ids")) \ .collect() source_row_dict = {row["source"]: row["row_ids"] for row in source_row_map} # 3. 迭代合并连通的行ID组 def merge_connected_groups(source_dict): groups = [] visited_rows = set() for rows in source_dict.values(): current_rows = set(rows) if not visited_rows.isdisjoint(current_rows): # 合并到已有组 merged = False for idx, group in enumerate(groups): if group & current_rows: groups[idx] = group.union(current_rows) merged = True break if not merged: groups.append(current_rows) else: groups.append(current_rows) visited_rows.update(current_rows) # 去重重复组 unique_groups = [g for i, g in enumerate(groups) if g not in groups[:i]] return unique_groups connected_groups = merge_connected_groups(source_row_dict) # 4. 将连通组转换为DataFrame并关联原始数据 group_data = [(row_id, group_idx) for group_idx, group in enumerate(connected_groups) for row_id in group] group_df = spark.createDataFrame(group_data, ["row_id", "component"]) # 5. 分组聚合得到结果 result_df = df_with_id.join(group_df, on="row_id") \ .groupBy("component") \ .agg( F.array_distinct(F.flatten(F.collect_set("values"))).alias("values"), F.array_distinct(F.flatten(F.collect_set("sources"))).alias("sources") ) \ .drop("component") result_df.show(truncate=False)
内容的提问来源于stack exchange,提问作者danni114
相关产品推荐
相关产品推荐

