Scala Spark中基于单条过滤记录带Where子句的Join实现
基于单条过滤记录筛选Spark DataFrame
原始数据与预处理
val columns = Seq("language", "users_count", "time_window") val data = Seq(("Java", "20000", "2021-04-05"),("Java", "20000", "2021-08-05"), ("Python", "100000", "2021-05-05"), ("Scala", "3000", "2021-07-05"), ("Python", "3000", "2021-03-05")) val rdd = spark.sparkContext.parallelize(data) val dfFromRDD1 = rdd.toDF(columns: _*) val dfFromRDD2 = dfFromRDD1.withColumn("time_window", date_format(col("time_window"), "yyyy-MM-dd")) .orderBy(desc("time_window")) // 按时间从新到旧倒序排序
需求
要筛选出最近一条Python条目(倒序排序后的第一条Python记录)及该条目之前的所有记录(即时间不晚于这条Python的所有数据),最终要保留的是:
- Java 20000 2021-08-05
- Scala 3000 2021-07-05
- Python 100000 2021-05-05
你的尝试代码
val dfFromRDD3 = dfFromRDD2.withColumn("idx", monotonically_increasing_id()) val filterRow = dfFromRDD3.filter(dfFromRDD3("language") === "Python").limit(1) val result = dfFromRDD3.as("df").join( filterRow.as("filterRow"), Seq("idx"), "left_outer" ).where($"df.idx" <= filterRow("idx").as[Integer])
问题修正与可行方案
你的思路方向没问题,但存在两个小问题:一是用left_outer join完全没必要,我们只需要拿到目标记录的阈值(索引或时间)来筛选;二是monotonically_increasing_id()生成的ID不保证连续,分布式环境下可能跳变,用它做索引风险高。
方案1:用连续行号筛选
先生成按排序顺序的连续行号,再提取目标行号做筛选:
import org.apache.spark.sql.expressions.Window // 生成按时间倒序排列的连续行号 val dfWithRowNum = dfFromRDD2.withColumn("row_num", row_number().over(Window.orderBy(desc("time_window")))) // 提取最近一条Python记录的行号 val targetRowNum = dfWithRowNum.filter($"language" === "Python") .select($"row_num") .limit(1) .first() .getAs[Int](0) // 筛选行号小于等于目标值的记录 val result = dfWithRowNum.filter($"row_num" <= targetRowNum)
方案2:直接用时间阈值筛选(更简洁)
既然需求是保留不早于最近Python记录的所有数据,直接提取这条记录的时间作为阈值即可:
// 提取最近一条Python记录的时间 val targetTime = dfFromRDD2.filter($"language" === "Python") .select($"time_window") .limit(1) .first() .getAs[String](0) // 筛选时间大于等于目标时间的记录(因为倒序排序,时间越大越新) val result = dfFromRDD2.filter($"time_window" >= targetTime)
说明
- 方案2更简洁高效,不需要额外生成索引列,完全基于业务字段筛选,更适配Spark分布式处理逻辑
- 方案1适合无法用业务字段直接筛选的场景,比如排序规则复杂、没有可用的时间/数值字段做阈值的情况
内容的提问来源于stack exchange,提问作者DataBach
相关产品推荐
相关产品推荐

