【问题标题】:Joining two Spark DataFrame according to size of intersection of two array columns根据两个数组列的交集大小加入两个 Spark DataFrame
【发布时间】:2016-11-26 14:23:32
【问题描述】:

我的 spark (v1.5.0) 代码中有两个 DataFrame

aDF = [user_id : Int, user_purchases: array<int> ]
bDF = [user_id : Int, user_purchases: array<int> ]

我想要做的是加入这两个数据框,但我只需要aDF.user_purchasesbDF.user_purchases 之间的交集有超过2个元素(交集> 2)的线。

我必须使用 RDD API 还是可以使用 org.apache.sql.functions 中的某些函数?

【问题讨论】:

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


    【解决方案1】:

    我没有看到任何内置函数,但你可以使用 UDF:

    import scala.collection.mutable.WrappedArray;
    val intersect = udf ((a : WrappedArray[Int], b : WrappedArray[Int]) => {
       var count = 0;
       a.foreach (x => {
           if (b.contains(x)) count = count + 1;
        });
        count;
    });
    // test data sets
    val one = sc.parallelize(List(
            (1, Array(1, 2, 3)), 
            (2, Array(1,2 ,3, 4)), 
            (3, Array(1, 2,3)), 
            (4, Array(1,2))
            )).toDF("user", "arr");
    
    val two = sc.parallelize(List(
            (1, Array(1, 2, 3)), 
            (2, Array(1,2 ,3, 4)), 
            (3, Array(1, 2, 3)), 
            (4, Array(1))
            )).toDF("user", "arr");
    
    // usage
    one.join(two, one("user") === two("user"))
        .select (one("user"), intersect(one("arr"), two("arr")).as("intersect"))
        .where(col("intersect") > 2).show
    
    // version from comment
    one.join(two)
        .select (one("user"), two("user"), intersect(one("arr"), two("arr")).as("intersect")).
        where('intersect > 2).show
    

    【讨论】:

    • 你的 udf 解决方案似乎解决了我的问题。只有一点注意,我不想加入相同的用户 id,我想要不同的用户 id,数组中至少有三个共同元素。
    • @Vektor88 所以在它之后进行交叉连接+过滤。会很慢,但没有其他选择
    • 非常感谢,我接受了您的回答,因为 udf 可以满足我的要求。您是否看到使用 intersect > 2 作为连接条件而不是稍后执行过滤器的任何副作用?我真的不需要将该值存储到列中
    • @Vektor88 你可以尝试运行解释 ;) 但据我所知,它将被重写为交叉连接 + 过滤器,但我目前没有安装任何 Spark 1.5 来检查它。 Spark 1.6 与 Spark 1.5 相比可能会有改进,所以我的explain 结果可能会有所不同
    【解决方案2】:

    一种可能的解决方案是找到有趣的配对并用数组扩充它们。首先让我们导入一些函数:

    import org.apache.spark.sql.functions.explode
    

    并重命名列:

    val aDF_ = aDF.toDF("a_user_id", "a_user_purchases")
    val bDF_ = bDF.toDF("b_user_id", "b_user_purchases")
    

    与谓词匹配的对可以被识别为:

    val filtered = aDF_.withColumn("purchase", explode($"a_user_purchases"))
      .join(bDF_.withColumn("purchase", explode($"b_user_purchases")), Seq("purchase"))
      .groupBy("a_user_id", "b_user_id")
      .count()
      .where($"count" > 2)
    

    最终过滤后的数据可以与输入数据集连接以获得完整结果:

    filtered.join(aDF_, Seq("a_user_id")).join(bDF_, Seq("b_user_id")).drop("count")
    

    在 Spark 2.4 或更高版本中,您还可以使用内置函数:

    import org.apache.spark.sql.functions.{size, array_intersect}
    
    aDF_
      .crossJoin(bDF_)
      .where(size(
        array_intersect($"a_user_purchases", $"b_user_purchases"
      )) > 2)
    

    虽然这可能仍然比更有针对性的哈希连接慢。

    【讨论】:

    • 聪明的 :) 但我认为它可能会很慢,因为物理计划要大得多并且包含很少的连接和排序
    猜你喜欢
    • 2017-07-16
    • 1970-01-01
    • 2018-12-19
    • 2019-06-04
    • 1970-01-01
    • 2019-06-06
    • 1970-01-01
    • 1970-01-01
    • 2020-07-12
    相关资源
    最近更新 更多