PySpark中按树形层级收集唯一id1列表的实现问题
问题背景
你有一个包含group_id、level、id1、id2的PySpark DataFrame,其中id2是子节点,id1是其父节点,数据形成树形层级结构。需求是为每个id2收集从自身向上到level=1的路径中的所有id1,要求每个层级仅保留对应路径的一个id1(例如id2=910657的路径应为[662867,677555,200001,0])。此前用窗口函数collect_list的方案会重复收集同层级id1且混入无关分支的id1,需要用PySpark API解决这个问题。
核心方案:递归CTE(推荐)
树形路径的本质是单节点到根节点的唯一链路,递归CTE(递归公共表表达式)能精准追踪每个节点的向上路径,避免无关数据混入。以下是具体实现步骤:
1. 样例数据准备
先创建模拟的树形结构DataFrame,方便演示:
from pyspark.sql import SparkSession from pyspark.sql.types import StructType, StructField, IntegerType spark = SparkSession.builder.appName("TreePathCollect").getOrCreate() data = [ (1, 4, 662867, 910657), (1, 3, 677555, 662867), (1, 2, 200001, 677555), (1, 1, 0, 200001), (1, 4, 123456, 789012), (1, 3, 234567, 123456), (1, 2, 200001, 234567) ] schema = StructType([ StructField("group_id", IntegerType(), True), StructField("level", IntegerType(), True), StructField("id1", IntegerType(), True), StructField("id2", IntegerType(), True) ]) df = spark.createDataFrame(data, schema) df.show()
2. 递归CTE实现路径收集
将DataFrame注册为临时表后,用递归CTE遍历每个节点的向上路径:
# 注册临时表供SQL查询使用 df.createOrReplaceTempView("tree_table") # 编写递归CTE查询语句 recursive_query = """ WITH RECURSIVE tree_path AS ( -- 基础项:初始化每个节点的路径为自身的父节点id1,记录当前节点id和层级 SELECT group_id, id2 AS current_id, level AS current_level, ARRAY(id1) AS path FROM tree_table UNION ALL -- 递归项:找到当前路径的父节点,将父节点的id1追加到路径中,直到层级为1 SELECT tp.group_id, tt.id2 AS current_id, tt.level AS current_level, array_union(tp.path, ARRAY(tt.id1)) AS path FROM tree_path tp JOIN tree_table tt ON tp.group_id = tt.group_id AND tp.path[0] = tt.id2 WHERE tt.level > 1 ) -- 提取每个原始节点的完整路径(按层级过滤,保留节点的最底层记录) SELECT current_id AS id2, group_id, path AS full_path FROM tree_path WHERE current_level = (SELECT MAX(level) FROM tree_table WHERE id2 = current_id) ORDER BY id2 """ # 执行查询并获取结果 result_df = spark.sql(recursive_query) result_df.show(truncate=False)
结果说明
执行后,id2=910657的full_path会是[662867,677555,200001,0],完全符合需求。递归过程中只会沿着当前节点的父节点向上遍历,不会混入其他分支的id1,也不会出现重复收集的问题。
为什么窗口函数方案失效?
窗口函数collect_list是按group_id分区、level排序后收集所有id1,无法区分树形结构中的分支关系。同一group_id下的不同分支节点会共享上层节点,但窗口函数会把整个分组内的id1全部收集,导致混入无关分支的节点,同时可能因为层级重复出现而收集到重复值。
备选方案:GraphFrames(需额外依赖)
如果需要更灵活的图操作,可以使用GraphFrames库构建树形图,然后通过最短路径功能获取节点到根的路径。需要先安装依赖:pip install graphframes,示例代码如下:
from graphframes import GraphFrame # 构建顶点表:所有节点id的去重集合 vertices = df.selectExpr("id1 as id").union(df.selectExpr("id2 as id")).distinct() # 构建边表:子节点(id2)指向父节点(id1) edges = df.selectExpr("id2 as src", "id1 as dst", "group_id") # 创建图对象 g = GraphFrame(vertices, edges) # 获取所有根节点(level=1的节点) root_nodes = df.filter(df.level == 1).select("id2").distinct().rdd.flatMap(lambda x: x).collect() # 计算每个节点到根节点的最短路径 path_result = g.shortestPaths(landmarks=root_nodes) # 关联原始DataFrame,整理出每个id2的路径 final_result = path_result.join(df, path_result.id == df.id2, "inner")\ .withColumn("full_path", path_result.distances[root_nodes[0]])\ .select("id2", "group_id", "full_path") final_result.show(truncate=False)
该方案适合复杂图场景,但递归CTE无需额外依赖,更适合大多数树形路径收集需求。
内容的提问来源于stack exchange,提问作者Henri

