如何用PySpark基于层级结构展开底层节点的记录?
PySpark层级结构底层节点路径生成方案
一、需求说明
给定层级结构的PySpark DataFrame,需要完成两个核心操作:
- 识别底层节点:即未出现在任何记录
ParentID字段中的节点(本例为C、D、Z) - 为每个底层节点生成完整层级路径的结构化数据:每条层级对应一条记录,ID字段为底层节点ID
二、原始数据定义
data = [("A",None, 1, 'Highest'),("B","A", 2, 'High'),("C","B", 3, 'Low'),("D","B", 3, 'Low'), ("X","A", 2, 'High'),("Y","X", 3, 'Low'),("Z","Y", 4, 'Lowest')] df = spark.createDataFrame(data=data, schema = ['ID','ParentID','Hierarchy','HierarchyName']) df.show(truncate=False)
原始数据输出:
+---+--------+---------+-------------+ |ID |ParentID|Hierarchy|HierarchyName| +---+--------+---------+-------------+ |A |null |1 |Highest | |B |A |2 |High | |C |B |3 |Low | |D |B |3 |Low | |X |A |2 |High | |Y |X |3 |Low | |Z |Y |4 |Lowest | +---+--------+---------+-------------+
三、完整实现步骤
1. 识别底层节点
通过anti join筛选出ID未出现在ParentID中的节点,这是最直接高效的方式:
from pyspark.sql import functions as F # 提取所有非空的父节点ID parent_ids = df.select("ParentID").distinct().filter(F.col("ParentID").isNotNull()) # 筛选出底层节点:ID不在父节点列表中的记录 leaf_nodes = df.select("ID").distinct().join(parent_ids, df.ID == parent_ids.ParentID, "anti")
执行后leaf_nodes结果:
+---+ |ID | +---+ |C | |D | |Z | +---+
2. 递归获取完整层级路径
通过循环递归向上遍历父节点,拼接每个底层节点的完整层级路径数组:
# 初始化递归DataFrame:仅保留底层节点的初始信息和路径 with_recursive = df.join(leaf_nodes, df.ID == leaf_nodes.ID, "inner") \ .withColumn("path", F.array(F.struct(F.col("Hierarchy"), F.col("HierarchyName")))) \ .withColumn("leaf_id", F.col("ID")) \ .select("leaf_id", "ID", "ParentID", "path") # 递归遍历父节点,将父节点层级信息拼接到路径前端 while True: temp = with_recursive.join(df, with_recursive.ParentID == df.ID, "left") \ .filter(df.ID.isNotNull()) \ .select( with_recursive.leaf_id, df.ID, df.ParentID, F.concat(F.array(F.struct(df.Hierarchy, df.HierarchyName)), with_recursive.path).alias("path") ) # 无更多父节点可关联时终止递归 if temp.count() == 0: break with_recursive = temp # 提取每个底层节点的完整路径 leaf_path_df = with_recursive.select("leaf_id", "path").distinct()
此时leaf_path_df结果(对应你思路中的df1):
+-------+---------------------------------------------------------------------+ |leaf_id|path | +-------+---------------------------------------------------------------------+ |C |[{3, Low}, {2, High}, {1, Highest}] | |D |[{3, Low}, {2, High}, {1, Highest}] | |Z |[{4, Lowest}, {3, Low}, {2, High}, {1, Highest}] | +-------+---------------------------------------------------------------------+
3. 展开路径生成目标结构
使用explode展开路径数组,提取层级信息并整理成目标格式:
# 展开路径数组 exploded_df = leaf_path_df.withColumn("hierarchy_item", F.explode(F.col("path"))) # 提取层级字段,并重命名ID为底层节点ID result_df = exploded_df.select( F.col("leaf_id").alias("ID"), F.col("hierarchy_item.Hierarchy").alias("Hierarchy"), F.col("hierarchy_item.HierarchyName").alias("HierarchyName") ) # 查看最终结果 result_df.show(truncate=False)
最终输出结果:
+---+---------+-------------+ |ID |Hierarchy|HierarchyName| +---+---------+-------------+ |C |3 |Low | |C |2 |High | |C |1 |Highest | |D |3 |Low | |D |2 |High | |D |1 |Highest | |Z |4 |Lowest | |Z |3 |Low | |Z |2 |High | |Z |1 |Highest | +---+---------+-------------+
四、完整可运行代码
from pyspark.sql import SparkSession from pyspark.sql import functions as F # 初始化SparkSession spark = SparkSession.builder.appName("HierarchyLeafPathGenerator").getOrCreate() # 1. 定义原始数据 data = [("A",None, 1, 'Highest'),("B","A", 2, 'High'),("C","B", 3, 'Low'),("D","B", 3, 'Low'), ("X","A", 2, 'High'),("Y","X", 3, 'Low'),("Z","Y", 4, 'Lowest')] df = spark.createDataFrame(data=data, schema = ['ID','ParentID','Hierarchy','HierarchyName']) # 2. 识别底层节点 parent_ids = df.select("ParentID").distinct().filter(F.col("ParentID").isNotNull()) leaf_nodes = df.select("ID").distinct().join(parent_ids, df.ID == parent_ids.ParentID, "anti") # 3. 递归获取完整层级路径 with_recursive = df.join(leaf_nodes, df.ID == leaf_nodes.ID, "inner") \ .withColumn("path", F.array(F.struct(F.col("Hierarchy"), F.col("HierarchyName")))) \ .withColumn("leaf_id", F.col("ID")) \ .select("leaf_id", "ID", "ParentID", "path") while True: temp = with_recursive.join(df, with_recursive.ParentID == df.ID, "left") \ .filter(df.ID.isNotNull()) \ .select( with_recursive.leaf_id, df.ID, df.ParentID, F.concat(F.array(F.struct(df.Hierarchy, df.HierarchyName)), with_recursive.path).alias("path") ) if temp.count() == 0: break with_recursive = temp leaf_path_df = with_recursive.select("leaf_id", "path").distinct() # 4. 展开路径生成目标结构 exploded_df = leaf_path_df.withColumn("hierarchy_item", F.explode(F.col("path"))) result_df = exploded_df.select( F.col("leaf_id").alias("ID"), F.col("hierarchy_item.Hierarchy").alias("Hierarchy"), F.col("hierarchy_item.HierarchyName").alias("HierarchyName") ) # 输出结果 result_df.show(truncate=False) # 停止SparkSession spark.stop()
内容的提问来源于stack exchange,提问作者tommyhmt
相关产品推荐
相关产品推荐

