Spark Scala:为DataFrame添加随ST段递增的ST_Group列
高效实现Spark DataFrame按ST段分组标记
核心思路
利用Spark的窗口函数实现分布式累计计数,完全规避本地循环的低效问题:
- 生成标识列,将
Segment = 'ST'的行标记为1,其余行标记为0; - 基于
Line_number的顺序,对标识列做累计求和,求和结果即为当前行所属的ST_Group编号。
完整Scala代码
import org.apache.spark.sql.expressions.Window import org.apache.spark.sql.functions.{col, lit, sum, when} // 1. 创建辅助标识列:标记ST起始行 val withStFlag = rawEDI.withColumn( "st_flag", when(col("Segment") === "ST", lit(1)).otherwise(lit(0)) ) // 2. 定义全局排序窗口:按Line_number保证行顺序 val windowSpec = Window.orderBy("Line_number") // 3. 累计求和生成分组编号 val resultDF = withStFlag .withColumn("ST_Group", sum(col("st_flag")).over(windowSpec)) .drop("st_flag") // 清理辅助列 // 输出结果 resultDF.show()
代码解释
- st_flag列:精准标记每个分组的起始点,只有ST行会触发计数递增;
- Window.orderBy("Line_number"):确保累计求和严格遵循EDI文件的行顺序,符合分组的自然逻辑;
- sum(col("st_flag")).over(windowSpec):从第一行到当前行的累计求和,每遇到一个ST,总和加1,后续行直到下一个ST前都会保持该数值,完美匹配分组编号需求。
原有方案问题分析
- 循环遍历:使用
collect()将全量数据拉取到Driver节点处理,彻底破坏Spark分布式特性;且代码未对每次循环结果做union,最终仅保留最后一次循环的过滤结果,导致返回空DataFrame; - 单纯UDF标记:仅标记ST行的位置,未实现累计和填充逻辑,无法为后续行分配相同组号;
- 错误Window分区:
partitionBy("Segment")会将相同Segment的行归为一组,row_number()仅生成各Segment内部的序号,完全不符合跨Segment的分组需求。
验证结果
运行代码后得到的DataFrame与预期完全一致:
| Line_number | Segment | ST_Group |
|---|---|---|
| 1 | ST | 1 |
| 2 | BPT | 1 |
| 3 | SE | 1 |
| 4 | ST | 2 |
| 5 | BPT | 2 |
| 6 | N1 | 2 |
| 7 | SE | 2 |
| 8 | ST | 3 |
| 9 | PTD | 3 |
| 10 | SE | 3 |
内容的提问来源于stack exchange,提问作者Peter Sanders
相关产品推荐
相关产品推荐

