【问题标题】:Aggregating arrays element wise聚合数组元素明智
【发布时间】:2023-03-25 20:34:01
【问题描述】:

spark/scala 的新手。我想知道是否有一种简单的方法可以按列方式聚合 Array[Double]。这是一个例子:

c1   c2   c3
-------------------------
1     1   [1.0, 1.0, 3.4]
1     2   [1.0, 0,0, 4.3]
2     1   [0.0, 0.0, 0.0]
2     3   [1.2, 1.1, 1.1]

然后,在聚合后,我会以如下所示的表格结束:

c1   c3prime
-------------
1     [2.0, 1.0, 7.7]
2     [1.2, 1.1, 1.1]

现在正在查看 UDAF,但想知道我是否需要编写代码?

感谢您的考虑。

【问题讨论】:

    标签: arrays scala apache-spark dataframe aggregation


    【解决方案1】:

    假设 c3 的数组值大小相同,您可以通过如下的 UDF 对列元素求和:

    val df = Seq(
      (1, 1, Seq(1.0, 1.0, 3.4)),
      (1, 2, Seq(1.0, 0.0, 4.3)),
      (2, 1, Seq(0.0, 0.0, 0.0)),
      (2, 3, Seq(1.2, 1.1, 1.1))
    ).toDF("c1", "c2", "c3")
    
    def elementSum = udf(
      (a: Seq[Seq[Double]]) => {
        val zeroSeq = Seq.fill[Double](a(0).size)(0.0)
        a.foldLeft(zeroSeq)(
          (a, x) => (a zip x).map{ case (u, v) => u + v }
        )
      }
    )
    
    val df2 = df.groupBy("c1").agg(
      elementSum(collect_list("c3")).as("c3prime")
    )
    
    df2.show(truncate=false)
    // +---+-----------------------------+
    // |c1 |c3prime                      |
    // +---+-----------------------------+
    // |1  |[2.0, 1.0, 7.699999999999999]|
    // |2  |[1.2, 1.1, 1.1]              |
    // +---+-----------------------------+
    

    【讨论】:

    • 嗨,Leo,这太棒了。非常感谢。人为设计的示例完全按照您的指定和预期工作。
    【解决方案2】:

    这是一个没有 UDF 的。它利用了 Spark 的 Window 函数。不知道效率如何,因为它涉及多个groupBys

    df.show
    
    // +---+---+---------------+
    // | c1| c2|             c3|
    // +---+---+---------------+
    // |  1|  1|[1.0, 1.0, 3.4]|
    // |  1|  2|[1.0, 0.0, 4.3]|
    // |  2|  1|[0.0, 0.0, 0.0]|
    // |  2|  2|[1.2, 1.1, 1.1]|
    // +---+---+---------------+
    
    import org.apache.spark.sql.expressions.Window
    
    val window = Window.partitionBy($"c1", $"c2").orderBy($"c1", $"c2")
    
    df.withColumn("c3", explode($"c3") )
      .withColumn("rn", row_number() over window)
      .groupBy($"c1", $"rn").agg(sum($"c3").as("c3") )
      .orderBy($"c1", $"rn")
      .groupBy($"c1")
      .agg(collect_list($"c3").as("c3prime") ).show
    
    // +---+--------------------+
    // | c1|             c3prime|
    // +---+--------------------+
    // |  1|[2.0, 1.0, 7.6999...|
    // |  2|     [1.2, 1.1, 1.1]|
    // +---+--------------------+
    

    【讨论】:

    • 谢谢慈善家。这也是一个有趣的观点。我会玩这个看看它是否有帮助。
    【解决方案3】:

    您可以组合一些inbuilt functions,例如groupBy、agg、sum、array、alias(as)等,得到所需的最终dataframe。

    import org.apache.spark.sql.functions._
    df.groupBy("c1")
      .agg(sum($"c3"(0)).as("c3_1"), sum($"c3"(1)).as("c3_2"), sum($"c3"(2)).as("c3_3"))
      .select($"c1", array("c3_1","c3_2","c3_3").as("c3prime"))
    

    希望回答对你有帮助。

    【讨论】:

      猜你喜欢
      • 2018-05-13
      • 2013-07-11
      • 2017-03-12
      • 2021-08-04
      • 2021-10-02
      • 2021-08-23
      • 2018-08-30
      • 1970-01-01
      • 2022-07-29
      相关资源
      最近更新 更多