Scala可变Set处理Spark DataFrame集群失败,如何改用Dataset?
我们有如下DataFrame:
val similarityJoinResult = toBeMatched1 .similarityJoin(toBeMatched2, "name", 0.6) .select( $"datasetA.directorId", $"datasetB.companyId", $"distance", )
函数similarityJoin行为类似普通inner join,但基于相似度匹配元素。distance列是相似度估计值,值越小名称相似度越高,因此会生成大量匹配结果——比如全量数据中该DF有1000万行,而唯一director ID仅1万个,需要去重并筛选最佳匹配。
设计的算法步骤
- 按
distance升序排序DF; - 逐行评估布尔标记;
- 若当前行的company ID或director ID已在之前标记为true的行中出现,标记设为false,否则设为true;
- 删除标记为false的行。
示例代码
def selectRowsWithUniqueIds(df: DataFrame): DataFrame = { val col1Ids: mutable.Set[String] = mutable.Set() val col2Ids: mutable.Set[String] = mutable.Set() def processRow(row: Row): (String, String, Boolean) = { val id1 = row.getString(0) val id2 = row.getString(1) val takeRow = !(col1Ids.contains(id1) || col2Ids.contains(id2)) if (takeRow) { col1Ids.add(id1) col2Ids.add(id2) } (id1, id2, takeRow) } df.map(processRow) .toDF(df.columns(0), df.columns(1), "takeRow") .filter($"takeRow") .drop("takeRow") } similarityJoinResult .orderBy($"distance".asc) .transform(selectRowsWithUniqueIds)
测试用例
示例DF
val df = Seq( ("id1", "myId1", 1), ("id1", "myId2", 1), ("id1", "myId3", 1), ("id2", "myId1", 15), ("id2", "myId2", 4), ("id2", "myId3", 2), ("id3", "myId1", 16), ("id3", "myId2", 5), ("id3", "myId3", 3), ).toDF("directorId", "companyId", "distance")
正确结果
+----------+---------+ |directorId|companyId| +----------+---------+ |id1 |myId1 | |id2 |myId3 | |id3 |myId2 | +----------+---------+
注意:按单个ID列分组再按另一列分组会导致数据丢失,可参考示例DF。
但该方案在集群上无法运行(本地正常),启动前因类初始化错误失败,推测是Scala可变Set导致。有以下疑问:
- 即使无初始化错误,Spark各Worker会有独立ID集合,导致保留重复数据,该假设是否正确?
- 若问题1答案为是,添加
.coalesce(1)能否解决? - 改用Dataset/DataFrame替换Scala集合,如何高效动态更新并检查值是否存在?
使用版本:Scala 2.12.17,Spark 3.3.2。
问题解答
1. 关于Worker节点独立ID集合的假设
完全正确。Spark的map操作是分布式执行的,每个Worker节点会拿到分区数据并独立执行processRow函数,每个节点上的可变Set都是独立初始化的,互相不共享状态。这会导致不同节点可能同时保留同一个directorId或companyId的匹配行,最终结果出现重复的ID对,完全不符合预期的去重逻辑。
2. 添加.coalesce(1)是否可行?
从功能上看,coalesce(1)会把所有数据合并到一个分区,这样整个处理流程会在单个节点上串行执行,可变Set的状态能保持一致,本地运行的逻辑可以正常生效。但不推荐在生产环境使用:1000万行数据合并到一个分区会导致单节点内存压力过大,处理速度极慢,甚至引发OOM,完全失去Spark分布式计算的优势。
3. 基于Dataset/DataFrame的高效实现方案
要实现全局的“已使用ID”跟踪,需要利用Spark的分布式状态管理或迭代式处理,推荐以下两种方案:
方案一:窗口函数+迭代过滤(适合数据量中等场景)
核心思路是每次筛选出当前未被使用的ID对,然后更新已使用的ID集合,重复这个过程直到没有新数据被筛选出来:
import org.apache.spark.sql.functions._ import org.apache.spark.sql.types.StringType def selectBestMatches(df: DataFrame): DataFrame = { val sortedDf = df.orderBy("distance") var usedDirectors = spark.createDataset(Seq.empty[String]).toDF("used_director") var usedCompanies = spark.createDataset(Seq.empty[String]).toDF("used_company") var result = spark.emptyDataFrame var remaining = sortedDf while (!remaining.isEmpty) { // 筛选当前未被使用的directorId和companyId的行 val candidate = remaining .join(usedDirectors, remaining("directorId") === usedDirectors("used_director"), "left_anti") .join(usedCompanies, remaining("companyId") === usedCompanies("used_company"), "left_anti") .limit(1) // 取当前最优的一行 if (candidate.isEmpty) { remaining = spark.emptyDataFrame } else { // 将当前行加入结果 result = result.union(candidate) // 更新已使用的ID集合 val dirId = candidate.select("directorId").head().getString(0) val compId = candidate.select("companyId").head().getString(0) usedDirectors = usedDirectors.union(spark.createDataset(Seq(dirId)).toDF("used_director")) usedCompanies = usedCompanies.union(spark.createDataset(Seq(compId)).toDF("used_company")) // 从剩余数据中排除包含已使用ID的行 remaining = remaining .filter(!col("directorId").isin(usedDirectors.collect().map(_.getString(0)): _*)) .filter(!col("companyId").isin(usedCompanies.collect().map(_.getString(0)): _*)) } } result.select("directorId", "companyId") } // 使用示例 val finalResult = selectBestMatches(similarityJoinResult)
方案二:使用RDD的mapPartitionsWithState(适合大数据量场景)
利用Spark RDD的状态管理API,在全局范围内跟踪已使用的ID:
import org.apache.spark.streaming.{State, StateSpec} import org.apache.spark.sql.Row // 定义状态类型:(已使用的directorId集合, 已使用的companyId集合) type IdState = (Set[String], Set[String]) def updateState(row: Row, state: State[IdState]): Option[(String, String)] = { val dirId = row.getString(0) val compId = row.getString(1) val currentState = state.getOption.getOrElse((Set.empty[String], Set.empty[String])) if (!currentState._1.contains(dirId) && !currentState._2.contains(compId)) { // 更新状态 val newDirs = currentState._1 + dirId val newComps = currentState._2 + compId state.update((newDirs, newComps)) Some((dirId, compId)) } else { None } } // 转换为RDD并排序 val sortedRdd = similarityJoinResult.orderBy("distance").rdd // 使用mapPartitionsWithState处理 val stateSpec = StateSpec.function(updateState _) val resultRdd = sortedRdd.mapPartitions(iter => iter.map(row => (0, row))) // 用固定key保证全局单状态 .mapPartitionsWithState(stateSpec) .filter(_.isDefined) .map(_.get) // 转换回DataFrame val finalResult = spark.createDataFrame(resultRdd).toDF("directorId", "companyId")
方案说明
- 方案一基于DataFrame API,代码更易读,但每次迭代需要收集已使用ID到Driver,适合10万级以内的唯一ID场景;
- 方案二基于RDD状态API,状态保存在Worker节点,Driver仅做协调,适合百万级唯一ID的大数据场景;
- 两种方案都避免了可变集合的分布式状态问题,能保证全局唯一的ID匹配逻辑。
内容的提问来源于stack exchange,提问作者Svyatoslav

