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

Spark中如何对DataFrame分组并使用函数转换每个分组?

如何对Spark RelationalGroupedDataset的每个分组应用自定义转换函数

当你通过df.groupBy("timestamp", "energy_type")得到RelationalGroupedDataset后,可通过以下几种方式实现自定义函数对每个分组的转换:

1. 使用mapGroups处理分组

mapGroups接收一个自定义函数,输入为分组键((timestamp, energy_type)元组)和该分组的行迭代器,返回转换后的行迭代器,适合一个分组输出一行结果的场景。

Scala 示例

import org.apache.spark.sql.{Row, SparkSession}
import org.apache.spark.sql.types._

// 定义转换后的输出Schema,按需调整
val outputSchema = StructType(Seq(
  StructField("timestamp", TimestampType),
  StructField("energy_type", StringType),
  StructField("unique_site_ids", ArrayType(StringType)),
  StructField("distinct_meter_count", IntegerType)
))

// 自定义分组处理函数:注意迭代器只能消费一次,需先缓存分组数据
def processGroup(key: (java.sql.Timestamp, String), iter: Iterator[Row]): Iterator[Row] = {
  val (timestamp, energyType) = key
  val groupRows = iter.toList
  val uniqueSites = groupRows.map(_.getAs[String]("site_id")).toSet.toArray
  val distinctMeters = groupRows.map(_.getAs[String]("meter_id")).distinct.size
  Iterator(Row(timestamp, energyType, uniqueSites, distinctMeters))
}

// 应用mapGroups完成转换
val transformedDF = df.groupBy("timestamp", "energy_type")
  .mapGroups(processGroup)(outputSchema)

Python 示例

from pyspark.sql import Row
from pyspark.sql.types import StructType, StructField, TimestampType, StringType, ArrayType, IntegerType

# 定义输出Schema
output_schema = StructType([
    StructField("timestamp", TimestampType()),
    StructField("energy_type", StringType()),
    StructField("unique_site_ids", ArrayType(StringType())),
    StructField("distinct_meter_count", IntegerType())
])

def process_group(key, group_iter):
    timestamp, energy_type = key
    group_rows = list(group_iter)
    unique_site_ids = list({row.site_id for row in group_rows})
    distinct_meter_count = len({row.meter_id for row in group_rows})
    yield Row(timestamp=timestamp, energy_type=energy_type, unique_site_ids=unique_site_ids, distinct_meter_count=distinct_meter_count)

# 应用mapGroups
transformed_df = df.groupBy("timestamp", "energy_type") \
    .mapGroups(process_group, output_schema)

2. 使用flatMapGroups处理分组

如果需要从一个分组输出多行结果,可使用flatMapGroups,它的逻辑和mapGroups类似,但返回的迭代器可包含多个行。

Python 示例

def flatten_group(key, group_iter):
    timestamp, energy_type = key
    # 将分组内的每个site_id拆分为单独行,保留分组键
    for row in group_iter:
        yield Row(timestamp=timestamp, energy_type=energy_type, site_id=row.site_id)

# 输出Schema沿用所需字段的结构
output_schema = df.select("timestamp", "energy_type", "site_id").schema
transformed_df = df.groupBy("timestamp", "energy_type") \
    .flatMapGroups(flatten_group, output_schema)

3. 使用applyInPandas(Python)处理分组

若习惯用Pandas处理分组逻辑,可使用applyInPandas:它会将每个分组转为Pandas DataFrame,处理后返回Pandas DataFrame,Spark自动合并结果。

Python 示例

import pandas as pd

def process_pandas_group(pdf):
    # pdf为当前分组对应的Pandas DataFrame
    pdf['unique_site_count'] = pdf['site_id'].nunique()
    pdf['distinct_meter_count'] = pdf['meter_id'].nunique()
    # 返回处理后的结果,需保留分组键
    return pdf[['timestamp', 'energy_type', 'unique_site_count', 'distinct_meter_count']].drop_duplicates()

# 应用applyInPandas
transformed_df = df.groupBy("timestamp", "energy_type") \
    .applyInPandas(process_pandas_group, schema=output_schema)

注意事项

  • 自定义函数的输入输出必须与定义的Schema严格匹配,否则会抛出类型不兼容错误。
  • 优先使用Spark内置聚合函数(如collect_set、countDistinct),内置函数的性能远优于自定义分组函数;仅当内置函数无法满足需求时再使用自定义转换。
  • Scala中使用mapGroups时,迭代器只能被消费一次,需先将分组数据缓存到集合(如List)再处理。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 06:45:33