PySpark:仅对DataFrame指定列Pivot并保留非聚合列单一值
问题描述
现有如下Spark DataFrame:
import pyspark.sql.functions as F df = spark.createDataFrame( [ [1, 'AB', 12, '2022-01-01'] , [1, 'AA', 22, '2022-01-10'] , [1, 'AC', 11, '2022-01-11'] , [2, 'AB', 22, '2022-02-01'] , [2, 'AA', 28, '2022-02-10'] , [2, 'AC', 25, '2022-02-22'] ] , 'code: int, doc_type: string, amount: int, load_date: string' ) df = df.withColumn('load_date', F.to_date('load_date'))
需求是:对amount列按doc_type执行Pivot操作,同时保留每个code分组下的第一个load_date值。
尝试了以下代码,但得到的结果中每个doc_type都对应了各自的load_date,不符合预期:
( df.groupBy('code') .pivot('doc_type', ['AB', 'AA', 'AC']) .agg(F.sum('amount').alias('amnt'), F.first('load_date').alias('ldt')) .show() )
执行结果:
+----+-------+----------+-------+----------+-------+----------+ |code|AB_amnt| AB_ldt|AA_amnt| AA_ldt|AC_amnt| AC_ldt| +----+-------+----------+-------+----------+-------+----------+ | 1| 12|2022-01-01| 22|2022-01-10| 11|2022-01-11| | 2| 22|2022-02-01| 28|2022-02-10| 25|2022-02-22| +----+-------+----------+-------+----------+-------+----------+
预期结果是每个code只保留一个全局的load_date,如下:
( df.groupBy('code') .agg( F.sum(F.when(F.col('doc_type') == 'AB', F.col('amount'))).alias('AB_amnt') , F.sum(F.when(F.col('doc_type') == 'AA', F.col('amount'))).alias('AA_amnt') , F.sum(F.when(F.col('doc_type') == 'AC', F.col('amount'))).alias('AC_amnt') , F.first('load_date').alias('load_date') ) .show() )
预期输出:
+----+-------+-------+-------+----------+ |code|AB_amnt|AA_amnt|AC_amnt| load_date| +----+-------+-------+-------+----------+ | 1| 12| 22| 11|2022-01-01| | 2| 22| 28| 25|2022-02-01| +----+-------+-------+-------+----------+
当前使用Databricks 14.3 LTS(Spark 3.5.0),由于存在多个需要Pivot的列和非Pivot列,希望找到更简洁的实现方式。
解决方案
方法1:拆分分组逻辑,再关联结果
先单独分组获取非Pivot的聚合列(比如每个code的第一个load_date),再和Pivot后的结果进行关联。这种方式逻辑清晰,适合多列Pivot和多非Pivot列的场景:
# 第一步:获取非Pivot列的聚合结果 non_pivot_df = df.groupBy('code').agg(F.first('load_date').alias('load_date')) # 第二步:执行Pivot操作 pivot_df = df.groupBy('code')\ .pivot('doc_type', ['AB', 'AA', 'AC'])\ .agg(F.sum('amount').alias('amnt')) # 第三步:关联两个结果 result_df = pivot_df.join(non_pivot_df, on='code', how='inner') result_df.show()
方法2:窗口函数预填充全局值,再执行Pivot
先通过窗口函数把每个code的第一个load_date填充到同组所有行,再按code和load_date分组Pivot,最终每个code只会保留一个load_date值:
from pyspark.sql import Window # 窗口定义:按code分组,按load_date排序取第一个值 window = Window.partitionBy('code').orderBy('load_date') df_with_ldt = df.withColumn('load_date', F.first('load_date').over(window)) # 分组Pivot result_df = df_with_ldt.groupBy('code', 'load_date')\ .pivot('doc_type', ['AB', 'AA', 'AC'])\ .agg(F.sum('amount').alias('amnt')) result_df.show()
方法3:动态生成聚合表达式(适合多Pivot维度)
如果有大量doc_type值,手动写when条件太繁琐,可以通过列表推导式动态生成聚合表达式,避免重复代码:
doc_types = ['AB', 'AA', 'AC'] # 动态生成Pivot列的聚合表达式 pivot_exprs = [F.sum(F.when(F.col('doc_type') == dt, F.col('amount'))).alias(f'{dt}_amnt') for dt in doc_types] # 非Pivot列的聚合表达式 non_pivot_exprs = [F.first('load_date').alias('load_date')] # 执行聚合 result_df = df.groupBy('code').agg(*pivot_exprs, *non_pivot_exprs) result_df.show()
内容的提问来源于stack exchange,提问作者Dhruv
相关产品推荐
相关产品推荐

