You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 01:14:58