如何在PySpark中将嵌套数组列拆分为多列?
PySpark 数组列拆分并转为多列实现
原始PySpark DataFrame:
+--------------------------------------------------------------------------------+ |substitutes | +--------------------------------------------------------------------------------+ |[[1, 30981733]] | |[[1, 598319049], [2, 38453298], [3, 2007569845]] | |[[1, 10309216]] | |[[1, 730446111], [2, 617024811], [3, 665689309], [4, 883699488], [5, 159896736]]| |[[1, 10290923], [2, 33282357]] | |[[1, 102649381], [2, 10294853], [3, 10294854], [4, 44749181], [5, 35132896]] | |[[1, 10307642], [2, 10307636], [3, 15754215], [4, 45612359], [5, 10307635]] | |[[1, 43982130], [2, 15556050], [3, 15556051], [4, 11961012], [5, 16777263]] | |[[1, 849607426], [2, 185158834], [3, 11028011], [4, 10309801], [5, 11028010]] | |[[1, 21905160], [2, 21609422], [3, 21609417], [4, 20554612], [5, 20554601]] | +--------------------------------------------------------------------------------+
期望结果:
| substitutes_1 | substitutes_2 | substitutes_3 | substitutes_4 | substitutes_5 | |---------------|---------------|---------------|---------------|---------------| | 30981733 | | | | | | 598319049 | 38453298 | 2007569845 | | | | 10309216 | | | | | | 730446111 | 617024811 | 665689309 | 883699488 | 159896736 | | 10290923 | 33282357 | | | | | 102649381 | 10294853 | 10294854 | 44749181 | 35132896 | | 10307642 | 10307636 | 15754215 | 45612359 | 10307635 | | 43982130 | 15556050 | 15556051 | 11961012 | 16777263 | | 849607426 | 185158834 | 11028011 | 10309801 | 11028010 | | 21905160 | 21609422 | 21609417 | 20554612 | 20554601 |
实现步骤与代码
核心逻辑:展开数组 → 拆分元素 → 透视列 → 重命名列
1. 导入依赖函数
from pyspark.sql import functions as F
2. 完整实现代码
# 给每行添加唯一ID,避免透视时数据混淆 df_with_id = df.withColumn("row_id", F.monotonically_increasing_id()) # 展开数组为单独行 exploded_df = df_with_id.withColumn("sub", F.explode("substitutes")) # 拆分每个数组元素为序号和对应值 split_df = exploded_df.withColumn("idx", F.col("sub")[0].cast("int")) \ .withColumn("value", F.col("sub")[1].cast("int")) # 透视生成目标列结构,指定列顺序确保排列符合预期 pivoted_df = split_df.groupBy("row_id") \ .pivot("idx", [1, 2, 3, 4, 5]) \ .agg(F.first("value")) # 重命名列并移除临时ID final_df = pivoted_df.select( *[F.col(str(i)).alias(f"substitutes_{i}") for i in range(1, 6)] ).drop("row_id") # 查看结果 final_df.show()
简化版代码(合并步骤)
from pyspark.sql import functions as F final_df = df.withColumn("row_id", F.monotonically_increasing_id()) \ .withColumn("sub", F.explode("substitutes")) \ .withColumn("idx", F.col("sub")[0].cast("int")) \ .withColumn("value", F.col("sub")[1].cast("int")) \ .groupBy("row_id") \ .pivot("idx", [1,2,3,4,5]) \ .agg(F.first("value")) \ .select(*[F.col(str(i)).alias(f"substitutes_{i}") for i in range(1,6)]) \ .drop("row_id") final_df.show()
内容的提问来源于stack exchange,提问作者Shibu
相关产品推荐
相关产品推荐

