【问题标题】:Spark merge/combine arrays in groupBy/aggregate在 groupBy/aggregate 中 Spark 合并/组合数组
【发布时间】:2017-01-22 14:50:45
【问题描述】:

以下 Spark 代码正确演示了我想要做什么,并使用一个很小的演示数据集生成正确的输出。

当我在大量生产数据上运行相同类型的通用代码时,我遇到了运行时问题。 Spark 作业在我的集群上运行了大约 12 个小时,但失败了。

只看下面的代码,将每一行都分解掉似乎效率很低,只是将其合并回来。在给定的测试数据集中,array_value_1 中包含三个值和array_value_2 中的三个值的第四行将爆炸为 3*3 或九个爆炸行。

那么,在一个更大的数据集中,一行有五个这样的数组列,每列有十个值,会爆炸成 10^5 行吗?

查看提供的 Spark 函数,没有开箱即用的函数可以满足我的需求。我可以提供一个用户定义的函数。这样做有速度上的缺点吗?

val sparkSession = SparkSession.builder.
  master("local")
  .appName("merge list test")
  .getOrCreate()

val schema = StructType(
  StructField("category", IntegerType) ::
    StructField("array_value_1", ArrayType(StringType)) ::
    StructField("array_value_2", ArrayType(StringType)) ::
    Nil)

val rows = List(
  Row(1, List("a", "b"), List("u", "v")),
  Row(1, List("b", "c"), List("v", "w")),
  Row(2, List("c", "d"), List("w")),
  Row(2, List("c", "d", "e"), List("x", "y", "z"))
)

val df = sparkSession.createDataFrame(rows.asJava, schema)

val dfExploded = df.
  withColumn("scalar_1", explode(col("array_value_1"))).
  withColumn("scalar_2", explode(col("array_value_2")))

// This will output 19. 2*2 + 2*2 + 2*1 + 3*3 = 19
logger.info(s"dfExploded.count()=${dfExploded.count()}")

val dfOutput = dfExploded.groupBy("category").agg(
  collect_set("scalar_1").alias("combined_values_2"),
  collect_set("scalar_2").alias("combined_values_2"))

dfOutput.show()

【问题讨论】:

    标签: scala apache-spark apache-spark-sql


    【解决方案1】:

    explode 可能效率低下,但从根本上说,您尝试实施的操作非常昂贵。实际上,它只是另一个groupByKey,您在这里无能为力让它变得更好。由于您使用 Spark > 2.0,因此您可以直接 collect_list 并展平:

    import org.apache.spark.sql.functions.{collect_list, udf}
    
    val flatten_distinct = udf(
      (xs: Seq[Seq[String]]) => xs.flatten.distinct)
    
    df
      .groupBy("category")
      .agg(
        flatten_distinct(collect_list("array_value_1")), 
        flatten_distinct(collect_list("array_value_2"))
      )
    

    在 Spark >= 2.4 中,您可以将 udf 替换为内置函数的组合:

    import org.apache.spark.sql.functions.{array_distinct, flatten}
    
    val flatten_distinct = (array_distinct _) compose (flatten _)
    

    也可以使用custom Aggregator,但我怀疑其中任何一个都会产生巨大的影响。

    如果集合相对较大并且您预计会有大量重复项,您可以尝试将aggregateByKey 与可变集合一起使用:

    import scala.collection.mutable.{Set => MSet}
    
    val rdd = df
      .select($"category", struct($"array_value_1", $"array_value_2"))
      .as[(Int, (Seq[String], Seq[String]))]
      .rdd
    
    val agg = rdd
      .aggregateByKey((MSet[String](), MSet[String]()))( 
        {case ((accX, accY), (xs, ys)) => (accX ++= xs, accY ++ ys)},
        {case ((accX1, accY1), (accX2, accY2)) => (accX1 ++= accX2, accY1 ++ accY2)}
      )
      .mapValues { case (xs, ys) => (xs.toArray, ys.toArray) }
      .toDF
    

    【讨论】:

    • 第一个简单的 flatten udf 解决方案完全解决了这个问题。 Spark 从失败前大约需要 12 小时变为在 30 分钟内成功完成整个工作。观察 Spark 监视器 GUI,每个内部任务都在一分钟或更短的时间内运行和完成。感谢您对此的帮助。
    • 我很高兴听到这个消息,尽管我不得不承认我很惊讶。我期待一个小的改进,但没有什么如此令人印象深刻的。单个列表有多大?
    • 您为我节省了数小时的搜索时间...非常感谢!
    猜你喜欢
    • 2022-12-07
    • 1970-01-01
    • 2020-01-20
    • 2016-12-12
    • 2019-11-08
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多