You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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导致。有以下疑问:

  1. 即使无初始化错误,Spark各Worker会有独立ID集合,导致保留重复数据,该假设是否正确?
  2. 若问题1答案为是,添加.coalesce(1)能否解决?
  3. 改用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 21:07:46