求Spark 1.6下累加和为负时重置为0的实现方案
Spark 1.6 Scala实现累加和负数重置逻辑方案
原实现代码
val sconf = new SparkConf() .setAppName("TestPonderacion") .setMaster("local[*]") var sc = new SparkContext(sconf) val sqlContext = new HiveContext(sc) val schema = StructType(List( StructField("CLIENT", IntegerType, true), StructField("FAMILY", StringType, true), StructField("MOV_DATE", StringType, true), StructField("IMP_PONDERACION", DoubleType, true) )) val data = List( Row(20898871, "VAUL 10", "2023-01-01 00:00:00", 25.00), Row(20898871, "VAUL 10", "2023-02-03 00:00:00", 25.00), Row(20898871, "VAUL 10", "2023-02-04 00:00:00", 1250.00), Row(20898871, "VAUL 10", "2023-03-01 00:00:00", -750.00), Row(20898871, "VAUL 10", "2023-03-02 00:00:00", 25.00), Row(20898871, "VAUL 10", "2023-03-03 00:00:00", 25.00), Row(20898871, "VAUL 10", "2023-04-01 00:00:00", -750.00), Row(20898871, "VAUL 10", "2023-04-02 00:00:00", 25.00), Row(20898871, "VAUL 10", "2023-04-03 00:00:00", 25.00) ) val mov = sqlContext.createDataFrame(sc.parallelize(data), schema) val movPonderados = mov.withColumn("CUMULATIVE_SUM", sum("IMP_PONDERACION").over(Window.partitionBy("FAMILY", "CLIENT").orderBy("MOV_DATE")).cast(DecimalType(17, 2))) movPonderados.printSchema() movPonderados.show(false)
当前运行结果
|CLIENT |FAMILY |MOV_DATE |IMP_PONDERACION|CUMULATIVE_SUM| +--------+-------+-------------------+---------------+--------------+ |20898871|VAUL 10|2023-01-01 00:00:00|25.0 |25.00 | |20898871|VAUL 10|2023-02-03 00:00:00|25.0 |50.00 | |20898871|VAUL 10|2023-02-04 00:00:00|1250.0 |1300.00 | |20898871|VAUL 10|2023-03-01 00:00:00|-750.0 |550.00 | |20898871|VAUL 10|2023-03-02 00:00:00|25.0 |575.00 | |20898871|VAUL 10|2023-03-03 00:00:00|25.0 |600.00 | |20898871|VAUL 10|2023-04-01 00:00:00|-750.0 |-150.00 | |20898871|VAUL 10|2023-04-02 00:00:00|25.0 |-125.00 | |20898871|VAUL 10|2023-04-03 00:00:00|25.0 |-100.00 | +--------+-------+-------------------+---------------+--------------+
需求说明
需要实现累加和重置逻辑:当累加和(前值+当前值)为负数时,将当前累加和重置为0,后续从0开始重新累加,期望结果如下:
|CLIENT |FAMILY |MOV_DATE |IMP_PONDERACION|CUMULATIVE_SUM| +--------+-------+-------------------+---------------+--------------+ |20898871|VAUL 10|2023-01-01 00:00:00|25.0 |25.00 | |20898871|VAUL 10|2023-02-03 00:00:00|25.0 |50.00 | |20898871|VAUL 10|2023-02-04 00:00:00|1250.0 |1300.00 | |20898871|VAUL 10|2023-03-01 00:00:00|-750.0 |550.00 | |20898871|VAUL 10|2023-03-02 00:00:00|25.0 |575.00 | |20898871|VAUL 10|2023-03-03 00:00:00|25.0 |600.00 | |20898871|VAUL 10|2023-04-01 00:00:00|-750.0 |0.00 | <- 因600-750为负,重置为0 |20898871|VAUL 10|2023-04-02 00:00:00|25.0 |25.00 | <- 从0开始重新累加 |20898871|VAUL 10|2023-04-03 00:00:00|25.0 |50.00 | +--------+-------+-------------------+---------------+--------------+
Spark 1.6适配解决方案
由于Spark 1.6的DataFrame窗口函数无法直接引用前一行的计算结果,我们可以通过RDD的分区迭代处理来实现逻辑:
步骤说明
- 将原DataFrame按
CLIENT和FAMILY分区,按MOV_DATE排序,确保同一分组的数据按时间顺序在同一个分区内。 - 对每个分区的迭代器进行遍历,维护一个累计和变量,逐个计算符合要求的累加值:
- 第一个元素:累计和取当前值与0的最大值
- 后续元素:计算前一个累计和加当前值,若结果为负则重置为0,否则保留结果
- 将处理后的RDD转回DataFrame,保持原字段并新增
CUMULATIVE_SUM字段。
实现代码
import org.apache.spark.sql._ import org.apache.spark.sql.types._ import scala.collection.mutable.ArrayBuffer import scala.math.BigDecimal.RoundingMode // 原代码部分保持不变,直到创建mov DataFrame... // 定义处理每个分区的函数 def processPartition(iter: Iterator[Row]): Iterator[Row] = { var cumulativeSum = 0.0 val result = ArrayBuffer[Row]() if (iter.hasNext) { val firstRow = iter.next() val firstValue = firstRow.getAs[Double]("IMP_PONDERACION") cumulativeSum = math.max(firstValue, 0.0) // 保留两位小数 val roundedSum = BigDecimal(cumulativeSum).setScale(2, RoundingMode.HALF_UP).toDouble result += Row.fromSeq(firstRow.toSeq :+ roundedSum) while (iter.hasNext) { val currentRow = iter.next() val currentValue = currentRow.getAs[Double]("IMP_PONDERACION") cumulativeSum = math.max(cumulativeSum + currentValue, 0.0) val roundedCurrentSum = BigDecimal(cumulativeSum).setScale(2, RoundingMode.HALF_UP).toDouble result += Row.fromSeq(currentRow.toSeq :+ roundedCurrentSum) } } result.iterator } // 转换为RDD处理并转回DataFrame val processedRDD = mov .repartitionByRange(1, $"CLIENT", $"FAMILY") // 可根据数据量调整分区数,确保同组数据在一个分区 .orderBy($"CLIENT", $"FAMILY", $"MOV_DATE") .rdd .mapPartitions(processPartition) // 构建新的Schema val newSchema = schema.add(StructField("CUMULATIVE_SUM", DoubleType, true)) val movPonderados = sqlContext.createDataFrame(processedRDD, newSchema) // 输出结果 movPonderados.printSchema() movPonderados.show(false)
注意事项
repartitionByRange的分区数可根据实际数据量调整,避免单分区数据量过大导致性能问题。- 使用
BigDecimal做精度处理,确保结果保留两位小数,与需求一致。
内容的提问来源于stack exchange,提问作者Fran
相关产品推荐
相关产品推荐

