【问题标题】:Dot product in Spark ScalaSpark Scala 中的点积
【发布时间】:2020-03-07 05:57:31
【问题描述】:

我在 Spark Scala 中有两个数据帧,每个数据帧的第二列是一个数字数组

val data22= Seq((1,List(0.693147,0.6931471)),(2,List(0.69314, 0.0)),(3,List(0.0, 0.693147))).toDF("ID","tf_idf")
data22.show(truncate=false)

+---+---------------------+
|ID |tf_idf               |
+---+---------------------+
|1  |[0.693, 0.702]       |
|2  |[0.69314, 0.0]       |
|3  |[0.0, 0.693147]      |
+---+---------------------+



val data12= Seq((1,List(0.69314,0.6931471))).toDF("ID","tf_idf")
data12.show(truncate=false)

+---+--------------------+
|ID |tf_idf              |
+---+--------------------+
|1  |[0.693, 0.805]      |
+---+--------------------+

我需要在这两个数据框中执行行之间的点积。那就是我需要将data12 中的tf_idf 数组与data22tf_idf 的每一行相乘。

(例如:点积的第一行应该是这样的:0.693*0.693 + 0.702*0.805

第二行:0.69314*0.693 + 0.0*0.805

第三行:0.0*0.693 + 0.693147*0.805)

基本上我想要一些东西(比如矩阵乘法)data22*transpose(data12) 如果有人可以建议一种在 Spark Scala 中执行此操作的方法,我将不胜感激。

谢谢

【问题讨论】:

    标签: scala apache-spark dot-product


    【解决方案1】:

    Spark 2.4+ 版: 使用zip_withaggregate 等数组函数,让您的代码更简单。按照您的详细描述,我已将join 更改为crossJoin

    val data22= Seq((1,List(0.693147,0.6931471)),(2,List(0.69314, 0.0)),(3,List(0.0, 0.693147))).toDF("ID","tf_idf")
    val data12= Seq((1,List(0.693,0.805))).toDF("ID2","tf_idf2")
    
    val df = data22.crossJoin(data12).drop("ID2")
    df.withColumn("DotProduct", expr("aggregate(zip_with(tf_idf, tf_idf2, (x, y) -> x * y), 0D, (sum, x) -> sum + x)")).show(false)
    

    这是结果。

    +---+---------------------+--------------+-------------------+
    |ID |tf_idf               |tf_idf2       |DotProduct         |
    +---+---------------------+--------------+-------------------+
    |1  |[0.693147, 0.6931471]|[0.693, 0.805]|1.0383342865       |
    |2  |[0.69314, 0.0]       |[0.693, 0.805]|0.48034601999999993|
    |3  |[0.0, 0.693147]      |[0.693, 0.805]|0.557983335        |
    +---+---------------------+--------------+-------------------+
    

    【讨论】:

    • 您好,谢谢您的回答。但这不是我想要的。你可能误解了我的问题。我很抱歉造成混乱。我更新了问题,所以请看一下。
    • @Sam88,我已经修改了答案。
    【解决方案2】:

    解决方法如下:

    scala> val data22= Seq((1,List(0.693147,0.6931471)),(2,List(0.69314, 0.0)),(3,List(0.0, 0.693147))).toDF("ID","tf_idf")
    data22: org.apache.spark.sql.DataFrame = [ID: int, tf_idf: array<double>]
    
    scala> val data12= Seq((1,List(0.69314,0.6931471))).toDF("ID","tf_idf")
    data12: org.apache.spark.sql.DataFrame = [ID: int, tf_idf: array<double>]
    
    scala> val arrayDot = data12.take(1).map(row => (row.getAs[Int](0), row.getAs[WrappedArray[Double]](1).toSeq))
    arrayDot: Array[(Int, Seq[Double])] = Array((1,WrappedArray(0.69314, 0.6931471)))
    
    scala> val dotColumn = arrayDot(0)._2
    dotColumn: Seq[Double] = WrappedArray(0.69314, 0.6931471)
    
    scala> val dotUdf = udf((y: Seq[Double]) => y zip dotColumn map(z => z._1*z._2) reduce(_ + _))
    dotUdf: org.apache.spark.sql.expressions.UserDefinedFunction = UserDefinedFunction(<function1>,DoubleType,Some(List(ArrayType(DoubleType,false))))
    
    scala> data22.withColumn("dotProduct", dotUdf('tf_idf)).show
    +---+--------------------+-------------------+
    | ID|              tf_idf|         dotProduct|
    +---+--------------------+-------------------+
    |  1|[0.693147, 0.6931...|   0.96090081381841|
    |  2|      [0.69314, 0.0]|0.48044305959999994|
    |  3|     [0.0, 0.693147]|    0.4804528329237|
    +---+--------------------+-------------------+
    

    请注意,它将data12 中的tf_idf 数组与data22tf_idf 的每一行相乘。

    如果有帮助请告诉我!!

    【讨论】:

      猜你喜欢
      • 2020-01-26
      • 1970-01-01
      • 1970-01-01
      • 2019-07-15
      • 2022-01-23
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2018-06-01
      相关资源
      最近更新 更多