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

PySpark中按树形层级收集唯一id1列表的实现问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 18:15:56