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

如何在PySpark的pivot操作中使用自定义聚合函数?

PySpark中Pivot使用自定义聚合函数的实现方法

常规Pivot用法回顾

PySpark中使用pivot进行数据透视的常规代码如下:

from pyspark.sql import SparkSession
from pyspark.sql.functions import col

# 创建SparkSession
spark = SparkSession.builder.getOrCreate()

# 构造测试DataFrame
data = [("Product A", "Region 1", 100),
        ("Product A", "Region 1", 150),
        ("Product A", "Region 2", 200),
        ("Product A", "Region 2", 250),
        ("Product B", "Region 1", 300),
        ("Product B", "Region 1", 350),
        ("Product B", "Region 2", 400),
        ("Product B", "Region 2", 450)]

df = spark.createDataFrame(data, ["Product", "Region", "SalesAmount"])

# 执行pivot,使用内置sum聚合
pivot_df = df.groupBy("Product").pivot("Region").sum("SalesAmount")

# 查看结果
pivot_df.show()

自定义聚合函数实现特殊操作

当内置聚合函数无法满足需求时,你需要先定义自定义用户聚合函数(UDAF),再替换掉代码中的sum即可。以下是两种主流实现方式:

方式1:传统UDAF(兼容Spark 2.x及以上)

以计算销售额的**最大值与最小值之差(范围)**为例,实现自定义聚合函数:

from pyspark.sql import SparkSession
from pyspark.sql.functions import col, UserDefinedAggregateFunction
from pyspark.sql.types import StructType, StructField, IntegerType

# 创建SparkSession
spark = SparkSession.builder.getOrCreate()

# 构造测试DataFrame
data = [("Product A", "Region 1", 100),
        ("Product A", "Region 1", 150),
        ("Product A", "Region 2", 200),
        ("Product A", "Region 2", 250),
        ("Product B", "Region 1", 300),
        ("Product B", "Region 1", 350),
        ("Product B", "Region 2", 400),
        ("Product B", "Region 2", 450)]

df = spark.createDataFrame(data, ["Product", "Region", "SalesAmount"])

# 定义自定义聚合函数:计算销售额范围(max - min)
class SalesRangeUDAF(UserDefinedAggregateFunction):
    # 输入数据的类型
    def inputSchema(self):
        return StructType([StructField("sales", IntegerType())])
    
    # 中间缓存数据的类型:存储当前分组的min和max
    def bufferSchema(self):
        return StructType([
            StructField("min_sales", IntegerType()),
            StructField("max_sales", IntegerType())
        ])
    
    # 最终输出结果的类型
    def dataType(self):
        return IntegerType()
    
    # 是否确定性输出(相同输入必得到相同输出)
    def deterministic(self):
        return True
    
    # 初始化缓存数据
    def initialize(self, buffer):
        buffer[0] = None  # min初始化为None
        buffer[1] = None  # max初始化为None
    
    # 处理每条输入数据,更新缓存
    def update(self, buffer, input):
        sales = input[0]
        if buffer[0] is None or sales < buffer[0]:
            buffer[0] = sales
        if buffer[1] is None or sales > buffer[1]:
            buffer[1] = sales
    
    # 合并多个分区的缓存数据
    def merge(self, buffer1, buffer2):
        # 合并min值
        if buffer1[0] is None:
            buffer1[0] = buffer2[0]
        elif buffer2[0] is not None and buffer2[0] < buffer1[0]:
            buffer1[0] = buffer2[0]
        # 合并max值
        if buffer1[1] is None:
            buffer1[1] = buffer2[1]
        elif buffer2[1] is not None and buffer2[1] > buffer1[1]:
            buffer1[1] = buffer2[1]
    
    # 计算最终输出结果
    def evaluate(self, buffer):
        return buffer[1] - buffer[0] if buffer[0] is not None and buffer[1] is not None else 0

# 创建自定义聚合函数实例
sales_range = SalesRangeUDAF()

# 使用自定义聚合函数执行pivot
pivot_df = df.groupBy("Product").pivot("Region").agg(sales_range(col("SalesAmount")))

# 查看结果
pivot_df.show()

方式2:Pandas UDAF(Spark 3.0+推荐)

Spark 3.0及以上支持基于Pandas的UDAF,代码更简洁,性能更优:

from pyspark.sql import SparkSession
from pyspark.sql.functions import col, pandas_udf
from pyspark.sql.types import IntegerType
import pandas as pd

# 创建SparkSession
spark = SparkSession.builder.getOrCreate()

# 构造测试DataFrame
data = [("Product A", "Region 1", 100),
        ("Product A", "Region 1", 150),
        ("Product A", "Region 2", 200),
        ("Product A", "Region 2", 250),
        ("Product B", "Region 1", 300),
        ("Product B", "Region 1", 350),
        ("Product B", "Region 2", 400),
        ("Product B", "Region 2", 450)]

df = spark.createDataFrame(data, ["Product", "Region", "SalesAmount"])

# 定义Pandas UDAF:计算销售额范围
@pandas_udf(IntegerType())
def sales_range(sales: pd.Series) -> int:
    return sales.max() - sales.min()

# 使用自定义聚合函数执行pivot
pivot_df = df.groupBy("Product").pivot("Region").agg(sales_range(col("SalesAmount")))

# 查看结果
pivot_df.show()

关键说明

  • 自定义聚合函数需要根据业务需求实现对应的逻辑,核心是定义好输入、缓存、输出的数据类型,以及数据更新、合并、计算的逻辑。
  • Pandas UDAF更适合熟悉Pandas语法的开发者,基于Arrow传输数据,性能优于传统UDAF。

内容的提问来源于stack exchange,提问作者figs_and_nuts

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 13:07:40