如何在PySpark DataFrame中实现数组转置与合并
PySpark实现数组转置与聚合需求
没问题,我来帮你搞定这个把多行数组转置成二维数组的需求!下面我会给出两种可行的实现方式,适配不同版本的Spark。
先创建测试数据
首先我们先把你提供的示例DataFrame创建出来,方便后续测试:
from pyspark.sql import SparkSession from pyspark.sql import functions as F spark = SparkSession.builder.appName("array_transpose_demo").getOrCreate() # 初始化示例数据 sample_data = [(1, [10, 11, 12]), (2, [20, 21, 22]), (3, [30, 31, 32])] df = spark.createDataFrame(sample_data, schema=["id", "data"]) # 查看原始数据 df.show()
方法一:使用posexplode + 分组聚合(兼容Spark 2.x+)
这个方法通过展开数组元素、分组聚合的方式实现,适合低版本的Spark:
# 1. 将data数组按索引位置展开,得到每个元素的位置和对应值 exploded_df = df.select(F.col("id"), F.posexplode(F.col("data")).alias("pos", "val")) # 2. 按索引位置分组,收集每个位置下的所有元素,形成单个数组 grouped_df = exploded_df.groupBy("pos").agg(F.collect_list("val").alias("column_vals")) # 3. 收集所有位置的数组形成二维数组,同时收集所有id组成数组,最后合并结果 result_df = ( grouped_df.agg(F.collect_list("column_vals").alias("data")) .crossJoin(df.agg(F.collect_list("id").alias("id"))) .select("id", "data") ) # 查看最终结果 result_df.show(truncate=False)
方法二:使用transform + aggregate(Spark 3.0+推荐)
如果你的Spark版本在3.0及以上,可以用更简洁的函数式写法,避免分组操作:
# 1. 先收集所有的id和所有的data数组 collected_df = df.agg( F.collect_list("id").alias("id"), F.collect_list("data").alias("all_data_arrays") ) # 2. 对收集到的data数组进行转置:遍历每个索引位置,收集所有数组对应位置的元素 result_df = collected_df.withColumn( "data", F.transform( # 生成从0到数组长度-1的索引序列 F.sequence(F.lit(0), F.size(F.col("all_data_arrays")[0]) - 1), # 对每个索引,聚合所有数组该位置的元素成新数组 lambda idx: F.aggregate( F.col("all_data_arrays"), F.array().cast("array<int>"), lambda acc, arr: F.concat(acc, F.array(arr[idx])) ) ) ).drop("all_data_arrays") # 查看最终结果 result_df.show(truncate=False)
两种方法最终都会输出你期望的结果:
+---------+----------------------------------+ |id |data | +---------+----------------------------------+ |[1, 2, 3]|[[10,20,30], [11,21,31], [12,22,32]]| +---------+----------------------------------+
内容的提问来源于stack exchange,提问作者jmlero
相关产品推荐
相关产品推荐

