PySpark遍历大型DataFrame替代方案:层级数据关联优化
问题描述
现有三个PySpark DataFrame:df_hrrchy、df_trans、df_creds,需关联生成指定结构的结果表。当前采用遍历df_trans的方式处理,因df_trans包含数百万条数据,效率极低,寻求无需遍历的高效实现方案。
输入数据示例
df_hrrchy
| lefId | Lineage |
|---|---|
| 36326 | ["36326","36465","36976","36091","82"] |
| 36121 | ["36121","36908","36976","36091","82"] |
| 36380 | ["36380","36465","36976","36091","82"] |
| 36448 | ["36448","36465","36976","36091","82"] |
| 36683 | ["36683","36465","36976","36091","82"] |
| 36949 | ["36949","36908","36976","36091","82"] |
| 37349 | ["37349","36908","36976","36091","82"] |
| 37026 | ["37026","36908","36976","36091","82"] |
| 36879 | ["36879","36465","36976","36091","82"] |
df_trans
| tranID | T_Id |
|---|---|
| 1000540 | ["36121","36326","37349","36949","36380","37026","36448","36683","36879"] |
df_creds
| T_Id | T_val | T_Goal | Parent_T_Id | Parent_Val | parent_Goal |
|---|---|---|---|---|---|
| 36448 | 100 | 1 | 36465 | 200 | 1 |
| 36465 | 200 | 1 | 36976 | 300 | 2 |
| 36326 | 90 | 1 | 36465 | 200 | 1 |
| 36091 | 500 | 19 | 82 | 600 | 4 |
| 36121 | 90 | 1 | 36908 | 200 | 1 |
| 36683 | 90 | 1 | 36465 | 200 | 1 |
| 36908 | 200 | 1 | 36976 | 300 | 2 |
| 36949 | 90 | 1 | 36908 | 200 | 1 |
| 36976 | 300 | 2 | 36091 | 500 | 19 |
| 37026 | 90 | 1 | 36908 | 200 | 1 |
| 37349 | 100 | 1 | 36908 | 200 | 1 |
| 36879 | 90 | 1 | 36465 | 200 | 1 |
| 36380 | 90 | 1 | 36465 | 200 | 1 |
期望结果
| T_id | children | T_Val | T_Goal | parent_T_id | parent_Goal | trans_id |
|---|---|---|---|---|---|---|
| 36091 | ["36976"] | 500 | 19 | 82 | 4 | 1000540 |
| 36465 | ["36448","36326","36683","36879","36380"] | 200 | 1 | 36976 | 2 | 1000540 |
| 36908 | ["36121","36949","37026","37349"] | 200 | 1 | 36976 | 2 | 1000540 |
| 36976 | ["36465","36908"] | 300 | 2 | 36091 | 19 | 1000540 |
| 36683 | null | 90 | 1 | 36465 | 1 | 1000540 |
| 37026 | null | 90 | 1 | 36908 | 1 | 1000540 |
| 36448 | null | 100 | 1 | 36465 | 1 | 1000540 |
| 36949 | null | 90 | 1 | 36908 | 1 | 1000540 |
| 36326 | null | 90 | 1 | 36465 | 1 | 1000540 |
| 36380 | null | 90 | 1 | 36465 | 1 | 1000540 |
| 36879 | null | 90 | 1 | 36465 | 1 | 1000540 |
| 36121 | null | 90 | 1 | 36908 | 1 | 1000540 |
| 37349 | null | 100 | 1 | 36908 | 1 | 1000540 |
尝试代码
from pyspark.sql import functions as F from pyspark.sql import DataFrame from pyspark.sql.functions import explode, collect_set, expr, col, collect_list,array_contains, lit from functools import reduce for row in df_transactions.rdd.toLocalIterator(): # def find_nodemap(row): dfs = [] df_hy_set = (df_hrrchy.filter(df_hrrchy.lefId.isin(row["T_ds"])) .select(explode("Lineage").alias("Terrs")) .agg(collect_set(col("Terrs")).alias("hierarchy_list")) .select(F.lit(row["trans_id"]).alias("trans_id"),"hierarchy_list") ) df_childrens = (df_creds.join(df_hy_set, expr("array_contains(hierarchy_list, T_id)")) .select("T_id", "T_Val","T_Goal","parent_T_id", "parent_Goal", "trans_id" ) .groupBy("parent_T_id").agg(collect_list("T_id").alias("children")) ) df_filter_creds = (df_creds.join(df_hy_set, expr("array_contains(hierarchy_list, T_id)")) .select ("T_id", "T_val","T_Goal","parent_T_id", "parent_Goal", "trans_id") ) df_nodemap = (df_filter_creds.alias("A").join(df_childrens.alias("B"), col("A.T_id") == col("B.parent_T_id"), "left") .select("A.T_id","B.children", "A.T_val","A.T_Goal","A.parent_T_id", "A.parent_Goal", "A.trans_id") ) display(df_nodemap) # dfs.append(df_nodemap) # df = reduce(DataFrame.union, dfs) # display(df) # # display(df)
当前设计不合理,df_trans包含数百万条数据,遍历DataFrame耗时极长,能否不通过遍历实现?尝试了其他方法,但未能得到期望结果。
解决方案
可以通过Spark的分布式关联操作替代遍历,核心思路是:
- 先将
df_trans中的T_Id数组展开,关联df_hrrchy获取对应的完整 lineage,再聚合每个tranID对应的所有层级节点集合。 - 基于层级节点集合关联
df_creds,并按tranID和parent_T_id分组统计子节点列表。 - 最后将子节点列表与
df_creds数据关联,得到最终结果。
具体代码如下:
from pyspark.sql import functions as F from pyspark.sql.functions import explode, collect_set, collect_list # 步骤1:展开df_trans的T_Id,关联df_hrrchy获取每个tranID对应的所有层级节点 df_trans_exploded = df_trans.withColumn("T_id_item", explode(F.col("T_Id"))) # 关联df_hrrchy获取lineage,再展开lineage得到所有节点 df_tran_hierarchy = df_trans_exploded.join( df_hrrchy, df_trans_exploded.T_id_item == df_hrrchy.lefId, "inner" ).withColumn("hierarchy_node", explode(F.col("Lineage"))) # 聚合每个tranID对应的所有层级节点集合 df_tran_nodes = df_tran_hierarchy.groupBy("tranID").agg( collect_set("hierarchy_node").alias("hierarchy_list") ) # 步骤2:关联df_creds,过滤出属于当前tranID层级的节点,并统计每个parent_T_id对应的子节点 df_children = df_tran_nodes.join( df_creds, F.expr("array_contains(hierarchy_list, T_Id)"), "inner" ).groupBy("tranID", "Parent_T_Id").agg( collect_list("T_Id").alias("children") ).withColumnRenamed("Parent_T_Id", "T_id") # 步骤3:关联df_creds和子节点列表,得到最终结果 df_final = df_creds.join( df_tran_nodes, F.expr("array_contains(hierarchy_list, T_Id)"), "inner" ).join( df_children, ["tranID", "T_id"], "left" ).select( F.col("T_Id").alias("T_id"), "children", "T_val", "T_Goal", F.col("Parent_T_Id").alias("parent_T_id"), F.col("parent_Goal").alias("parent_Goal"), F.col("tranID").alias("trans_id") ) # 展示结果 df_final.show(truncate=False)
代码说明
- 步骤1:将每个交易的
T_Id数组展开,关联df_hrrchy获取每个初始节点的完整 lineage,再将lineage展开后聚合,得到每个交易对应的所有层级节点集合,避免逐行遍历。 - 步骤2:通过数组包含判断过滤出每个交易的层级节点,再按交易ID和父节点分组,收集子节点列表。
- 步骤3:将
df_creds与交易节点集合关联,再左连接子节点列表,最终整理成期望的字段结构。
这种方式完全利用Spark的分布式计算能力,无需遍历单条交易数据,能高效处理数百万条df_trans记录。
内容的提问来源于stack exchange,提问作者sys
相关产品推荐
相关产品推荐

