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

Spark DataFrame依据另一数组列做聚合的性能优化问题

Spark 2.4下按索引数组聚合双精度数组的高性能实现方案

需求说明

需要按id分组,将同id下所有prop1(索引数组)对应的values(值数组)做同索引值求和,最终输出每个id对应的去重索引数组,以及对应位置的聚合值数组。
原始数据结构示例:

| id | prop1        | values                  |
|----|--------------|-------------------------|
|  1 | [2,5,1,3]    |   [ 0.1, 0.5, 0.7, 0.8] |
|  2 | [2,1]        |   [ 0.2, 0.3 ]          |
|  1 | [1,5]        |   [ 0.4, 0.3 ]          |
|  2 | [3,2]        |   [ 0.0, 0.1 ]          |

预期输出:

| id | prop1          | values                   | 
|----|----------------|--------------------------| 
|  1 | [2,5,1,3]      |   [ 0.1, 0.8, 1.1, 0.8 ] | 
|  2 | [2,1,3]        |   [ 0.3, 0.3, 0.0 ]      |

现有方案性能问题根源

当前使用的explode+join+pivot方案在大数据量下失效的核心原因:

  • 大数组explode会导致数据量膨胀数万倍,单个数组长度50万的行直接变成50万行,内存和IO开销爆炸
  • pivot操作当prop1取值上限达到30万时会生成数十万列,Spark根本无法处理这么宽的表

优化方案(完全避免数据膨胀)

核心思路:全程基于数组/Map做行内运算,不做任何展开操作,通过自定义UDAF合并同id的键值对Map完成聚合,Spark 2.4可完美支持。

实现步骤

  1. 定义合并Map的自定义UDAF:功能是输入多个Map[Int, Double],输出合并后的Map,相同key的value做累加
  2. 每行将prop1和values转成Map[Int, Double]
  3. 按id分组,用自定义UDAF聚合所有Map
  4. 从聚合后的Map中提取key、value数组即为最终结果

代码示例

import org.apache.spark.sql.expressions.MutableAggregationBuffer
import org.apache.spark.sql.expressions.UserDefinedAggregateFunction
import org.apache.spark.sql.types._
import org.apache.spark.sql.functions._

// 自定义Map合并求和UDAF
object MapSumUDAF extends UserDefinedAggregateFunction {
  // 输入数据类型:Map[Int, Double]
  override def inputSchema: StructType = StructType(StructField("inputMap", MapType(IntegerType, DoubleType)) :: Nil)

  // 中间缓冲数据类型:Map[Int, Double]
  override def bufferSchema: StructType = StructType(StructField("bufferMap", MapType(IntegerType, DoubleType)) :: Nil)

  // 输出数据类型:Map[Int, Double]
  override def dataType: DataType = MapType(IntegerType, DoubleType)

  override def deterministic: Boolean = true

  // 初始化缓冲
  override def initialize(buffer: MutableAggregationBuffer): Unit = {
    buffer(0) = Map.empty[Int, Double]
  }

  // 单个分区内合并新的输入Map到缓冲
  override def update(buffer: MutableAggregationBuffer, input: org.apache.spark.sql.Row): Unit = {
    val inputMap = input.getAs[Map[Int, Double]](0)
    val bufferMap = buffer.getAs[Map[Int, Double]](0)
    buffer(0) = bufferMap ++ inputMap.map { case (k, v) => k -> (bufferMap.getOrElse(k, 0.0) + v) }
  }

  // 跨分区合并缓冲
  override def merge(buffer1: MutableAggregationBuffer, buffer2: org.apache.spark.sql.Row): Unit = {
    val map1 = buffer1.getAs[Map[Int, Double]](0)
    val map2 = buffer2.getAs[Map[Int, Double]](0)
    buffer1(0) = map1 ++ map2.map { case (k, v) => k -> (map1.getOrElse(k, 0.0) + v) }
  }

  // 返回最终结果
  override def evaluate(buffer: org.apache.spark.sql.Row): Any = {
    buffer.getAs[Map[Int, Double]](0)
  }
}

// 注册UDAF
spark.udf.register("map_sum", MapSumUDAF)

// 主逻辑
val result = originalDf
  // 每行转prop1->values的Map
  .withColumn("prop_val_map", map_from_arrays(col("prop1"), col("values")))
  // 按id分组聚合Map
  .groupBy("id")
  .agg(MapSumUDAF(col("prop_val_map")).alias("agg_map"))
  // 提取key、value数组为最终输出
  .select(
    col("id"),
    map_keys(col("agg_map")).alias("prop1"),
    map_values(col("agg_map")).alias("values")
  )

额外优化点

  • 如果单行的prop1数组存在重复索引,可在转Map前先对单行的索引和值做预处理,同索引先求和,避免转Map时值被覆盖
  • 可根据id的基数调整shuffle分区数,避免数据倾斜

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 07:15:04