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可完美支持。
实现步骤
- 定义合并Map的自定义UDAF:功能是输入多个
Map[Int, Double],输出合并后的Map,相同key的value做累加 - 每行将
prop1和values转成Map[Int, Double] - 按
id分组,用自定义UDAF聚合所有Map - 从聚合后的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
相关产品推荐
相关产品推荐

