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

如何基于列中部分重叠值合并PySpark DataFrame行?

基于PySpark DataFrame合并共享来源的行

需求:当DataFrame的行之间共享sources列中的值时,合并values集合与去重后的sources集合;无共享来源的行保持不变。已完成sources列的explode操作,寻求后续实现方法。

原DataFrame示例

valuessources
[a, b][s1, s2, s3]
[b, c, d][s5, s1]
[x, y][s7]

期望输出

valuessources
[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

valuessource
[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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 02:06:01