Scala Spark数组列:计算1与2间连续0的最大出现次数
问题:Spark DataFrame中统计1和2之间连续0的最大出现次数
现有一个Scala Spark DataFrame,其Schema如下:
root |-- passengerId: string (nullable = true) |-- travelHist: array (nullable = true) | |-- element: integer (containsNull = true)
需求:遍历数组元素,找出位于1和2之间的连续0的最大出现次数(仅统计1之后、2之前的连续0序列)。
输入示例
| passengerID | travelHist |
|---|---|
| 1 | 1, 0, 0, 0, 0, 2, 1, 0, 0, 0, 0, 0, 0, 0, 2, 1, 0 |
| 2 | 0, 0, 0, 0, 0, 0, 0, 0, 2, 1, 0, 0, 0, 2, 0, 0, 0, 0 |
| 3 | 0,0,0,2,1,0,2,1,0 |
预期输出
| passengerID | maxStreak |
|---|---|
| 1 | 7 |
| 2 | 3 |
| 3 | 1 |
假设数组元素数量不超过50个,请问实现该需求的最高效方式是什么?
实现方案
因为数组长度限制在50以内,优先使用Spark内置高阶函数实现,避免自定义UDF的序列化开销,同时保证代码简洁高效。核心思路是定位1和2的有效区间,仅统计这些区间内的连续0长度,最后取最大值。
具体Scala代码实现
import org.apache.spark.sql.functions._ import org.apache.spark.sql.types._ // 模拟输入DataFrame val inputDF = spark.createDataFrame(Seq( ("1", Array(1,0,0,0,0,2,1,0,0,0,0,0,0,0,2,1,0)), ("2", Array(0,0,0,0,0,0,0,0,2,1,0,0,0,2,0,0,0,0)), ("3", Array(0,0,0,2,1,0,2,1,0)) )).toDF("passengerId", "travelHist") // 核心逻辑 val resultDF = inputDF // 给数组元素添加位置索引 .withColumn("indexed_hist", transform($"travelHist", (value, idx) => struct(idx as "pos", value as "val"))) // 提取所有值为1或2的标记点 .withColumn("key_points", filter($"indexed_hist", x => x("val") === 1 || x("val") === 2)) // 生成连续标记点的区间(前一个标记到后一个标记) .withColumn("valid_ranges", transform( slice($"key_points", 1, size($"key_points") - 1), (x, idx) => struct( $"key_points"(idx)("pos") + 1 as "start", x("pos") - 1 as "end", $"key_points"(idx)("val") as "left_val", x("val") as "right_val" ) )) // 过滤出左值为1、右值为2的有效区间 .withColumn("valid_ranges", filter($"valid_ranges", x => x("left_val") === 1 && x("right_val") === 2)) // 统计每个有效区间内连续0的最大长度 .withColumn("streaks", transform($"valid_ranges", range => { val subArr = slice($"travelHist", range("start") + 1, range("end") - range("start") + 1) aggregate( subArr, struct(lit(0) as "current", lit(0) as "max"), (acc, elem) => when(elem === 0, struct(acc("current") + 1 as "current", greatest(acc("max"), acc("current") + 1) as "max")) .otherwise(struct(lit(0) as "current", acc("max") as "max")), acc => acc("max") ) })) // 取所有区间的最大连续0长度,无有效区间则返回0 .withColumn("maxStreak", coalesce(array_max($"streaks"), lit(0))) // 保留目标列 .select($"passengerId", $"maxStreak") resultDF.show()
代码说明
- indexed_hist:给数组每个元素添加位置索引,方便后续定位区间范围
- key_points:提取所有
1和2的标记点,快速缩小需要处理的范围 - valid_ranges:筛选出
1后跟2的有效区间,确保只统计符合需求的片段 - streaks:对每个有效区间的子数组,用
aggregate函数实时统计连续0的长度,记录最大值 - maxStreak:取所有有效区间的最大值,没有符合条件的区间时返回0
效率优势
- 全程使用Spark内置高阶函数,避免了自定义UDF带来的序列化/反序列化开销,执行效率更高
- 针对数组长度≤50的特点,所有操作都是轻量级内存计算,无额外性能负担
- 仅处理有效区间,避免了遍历整个数组的冗余计算
内容的提问来源于stack exchange,提问作者Abishek
相关产品推荐
相关产品推荐

