Spark实现带阈值限制的累计求和(非UDAF方案求助)
带阈值限制的Spark分区内累计求和实现(非UDAF方案)
业务需求
在Spark数据集的分区内计算累计求和(running sum),要求累计求和结果不得超过指定阈值(threshold),核心规则:
- 同一
id的所有行阈值相同 - 若当前行加入后累计和超过阈值,则跳过该行,累计和保持上一行结果
- 若当前行加入后累计和未超阈值,则将该行值计入累计和
示例数据
输入数据集
| id | created_on | value | running_sum | threshold |
|---|---|---|---|---|
| A | 2021-01-01 | 1.0 | 0.0 | 10.0 |
| A | 2021-01-02 | 2.0 | 0.0 | 10.0 |
| A | 2021-01-03 | 8.0 | 0.0 | 10.0 |
| A | 2021-01-04 | 5.0 | 0.0 | 10.0 |
预期输出
| id | created_on | value | running_sum | threshold |
|---|---|---|---|---|
| A | 2021-01-01 | 1.0 | 1.0 | 10.0 |
| A | 2021-01-02 | 2.0 | 3.0 | 10.0 |
| A | 2021-01-03 | 8.0 | 3.0 | 10.0 |
| A | 2021-01-04 | 5.0 | 8.0 | 10.0 |
现有尝试的问题
- 无阈值限制的窗口函数实现:仅能计算普通累计和,无法处理阈值限制
final WindowSpec window = Window.partitionBy(col("id")) .orderBy(col("created_on").asc()) .rowsBetween(Window.unboundedPreceding(), Window.currentRow()); dataset.withColumn("running_sum", sum(col("value")).over(window)).show();
- sum+when组合尝试:第四行累计和计算错误(得到11.0),因为该方案仅判断全局累计和是否超阈值,未考虑上一行的有效累计值
final WindowSpec window = Window.partitionBy(col("id")) .orderBy(col("created_on").asc()) .rowsBetween(Window.unboundedPreceding(), Window.currentRow()); dataset.withColumn("running_sum", when(sum(col("value")).over(window).leq(col("threshold")), sum(col("value")).over(window)) .otherwise(sum(col("value")).over(window).minus(col("value"))) ).show();
- sum+lag组合尝试:第三、四行累计和为0.0,因为
running_sum是新生成列,lag函数无法引用同一语句中刚计算出的列值,导致逻辑失效
final WindowSpec lagWindow = Window.partitionBy(col("id")).orderBy(col("created_on").asc()); final WindowSpec window = Window.partitionBy(col("id")) .orderBy(col("created_on").asc()) .rowsBetween(Window.unboundedPreceding(), Window.currentRow()); dataset.withColumn("running_sum", when(sum(col("value")).over(window).leq(col("threshold")), sum(col("value")).over(window)) .otherwise(lag(col("running_sum"), 1, 0).over(lagWindow)) ).show();
非UDAF解决方案:使用aggregate函数结合窗口函数
Spark的aggregate函数可在窗口内实现自定义累加逻辑,满足带阈值的累计求和需求,以下是Java实现代码:
// 定义分区排序的窗口 final WindowSpec window = Window.partitionBy(col("id")) .orderBy(col("created_on").asc()); // 使用aggregate实现带阈值的累计求和 Dataset<Row> result = dataset.withColumn("running_sum", aggregate( // 收集窗口内从起始到当前行的value序列 collect_list(col("value")).over(window), // 初始累计值 lit(0.0), // 累加逻辑:判断累计值+当前value是否超阈值,是则保持原值,否则累加 (acc, val) -> when(acc.plus(val).leq(col("threshold")), acc.plus(val)).otherwise(acc), // 最终输出累计值 acc -> acc ) ); result.show();
逻辑说明
collect_list(col("value")).over(window):将窗口内从第一行到当前行的value收集为列表aggregate遍历列表,从初始值0.0开始累加:- 每次判断当前累计值加当前元素是否小于等于阈值
- 满足条件则更新累计值,否则保持原累计值
- 最终输出的累计值即为符合阈值要求的running sum
注意事项
- 确保
created_on的排序逻辑正确,保证累计求和顺序符合业务要求 - 若数据集较大,
collect_list可能占用较多内存,此时可考虑分批次处理,或评估UDAF的必要性(该方案已满足非UDAF需求)
内容的提问来源于stack exchange,提问作者sanketd617
相关产品推荐
相关产品推荐

