基于别名数组交集关联并合并多DataFrame的实现方案
问题需求
需要以别名数组为关联条件,连接三个Spark DataFrame,生成指定结构的结果表:将所有关联别名归为同一行,且每个表的ID不重复——本质是按连通别名组(共享同一ID关联的别名属于同一组)分组,收集每组的所有别名及各表中出现该组别名的唯一ID列表。支持Spark SQL、PySpark或Pandas实现。
数据源定义
Table 1
table_1 = spark.createDataFrame([("T1", ['a','b','c']), ("T2", ['d','e','f'])], ["id", "aliases"])
Table 2
table_2 = spark.createDataFrame([("P1", ['a','h','e']), ("P2", ['j','k','l'])], ["id", "aliases"])
Table 3
table_3 = spark.createDataFrame([("G1", ['a','n','o']), ("G2", ['p','q','l']), ("G3", ['c','z'])], ["id", "aliases"])
期望输出
| Aliases | table1_ids | table2_id | table3_id |
|---|---|---|---|
| [n, b, h, o, a, e, d, c, f, z] | [T1, T2] | [P1] | [G1,G3] |
| [k, q, j, p, l] | [] | [P2] | [G2] |
PySpark 实现方案
核心思路是:先拆分别名数组为单行记录,再通过连通分量算法识别关联的别名组,最后按组聚合结果。
方式1:使用GraphFrames(推荐,简洁高效)
需先确保环境安装graphframes:pip install graphframes
from pyspark.sql import functions as F from graphframes import GraphFrame # 1. 拆分各表的别名数组,保留来源ID t1_exploded = table_1.select(F.col("id").alias("source_id"), F.explode("aliases").alias("alias")) t2_exploded = table_2.select(F.col("id").alias("source_id"), F.explode("aliases").alias("alias")) t3_exploded = table_3.select(F.col("id").alias("source_id"), F.explode("aliases").alias("alias")) # 2. 合并所有别名与ID的映射关系 all_mappings = t1_exploded.union(t2_exploded).union(t3_exploded) # 3. 构建图结构:顶点为别名,边为同一ID下的别名关联 vertices = all_mappings.select("alias").distinct().withColumnRenamed("alias", "id") edges = all_mappings.alias("a").join(all_mappings.alias("b"), on="source_id") \ .select(F.col("a.alias").alias("src"), F.col("b.alias").alias("dst")) \ .filter(F.col("src") != F.col("dst")) # 4. 计算连通分量(识别关联别名组) graph = GraphFrame(vertices, edges) connected_components = graph.connectedComponents() # 5. 按连通组分组合并结果 result = all_mappings.join(connected_components, all_mappings.alias == connected_components.id, "left") \ .groupBy("component") \ .agg( F.collect_set("alias").alias("Aliases"), # 收集各表的唯一ID,过滤非对应表的ID F.collect_set(F.when(F.col("source_id").startswith("T"), F.col("source_id"))).alias("table1_ids"), F.collect_set(F.when(F.col("source_id").startswith("P"), F.col("source_id"))).alias("table2_id"), F.collect_set(F.when(F.col("source_id").startswith("G"), F.col("source_id"))).alias("table3_id") ) \ .select("Aliases", "table1_ids", "table2_id", "table3_id") # 展示结果 result.show(truncate=False)
方式2:无需GraphFrames,用窗口函数实现
如果无法使用GraphFrames,可通过递归窗口函数模拟连通分量计算:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 1. 拆分别名数组并标记来源表 t1_exploded = table_1.select(F.col("id").alias("id"), F.explode("aliases").alias("alias"), F.lit("t1").alias("src")) t2_exploded = table_2.select(F.col("id").alias("id"), F.explode("aliases").alias("alias"), F.lit("t2").alias("src")) t3_exploded = table_3.select(F.col("id").alias("id"), F.explode("aliases").alias("alias"), F.lit("t3").alias("src")) all_data = t1_exploded.union(t2_exploded).union(t3_exploded) # 2. 生成同一ID下的别名配对关系 alias_pairs = all_data.alias("a").join(all_data.alias("b"), on="id") \ .select(F.col("a.alias").alias("alias1"), F.col("b.alias").alias("alias2")) \ .filter(F.col("alias1") != F.col("alias2")) # 3. 迭代计算连通分量(传递闭包) connected = alias_pairs.withColumn("component", F.col("alias1")) for _ in range(5): # 迭代次数可根据数据复杂度调整 w = Window.partitionBy("alias2") connected = connected.alias("a") \ .join(connected.alias("b"), F.col("a.component") == F.col("b.alias1"), "left") \ .withColumn("new_component", F.coalesce(F.col("b.component"), F.col("a.component"))) \ .select("a.alias1", "a.alias2", "new_component") \ .withColumnRenamed("new_component", "component") # 确定每个别名的最终分组 alias_groups = connected.select("alias1", "component") \ .union(connected.select("alias2", "component")) \ .distinct() \ .withColumn("component", F.min("component").over(Window.partitionBy("alias1"))) # 4. 按分组聚合结果 final_result = all_data.join(alias_groups, all_data.alias == alias_groups.alias1, "left") \ .groupBy("component") \ .agg( F.collect_set("alias").alias("Aliases"), F.collect_set(F.when(F.col("src") == "t1", F.col("id"))).alias("table1_ids"), F.collect_set(F.when(F.col("src") == "t2", F.col("id"))).alias("table2_id"), F.collect_set(F.when(F.col("src") == "t3", F.col("id"))).alias("table3_id") ) \ .select("Aliases", "table1_ids", "table2_id", "table3_id") final_result.show(truncate=False)
内容的提问来源于stack exchange,提问作者AngryCoder
相关产品推荐
相关产品推荐

