【问题标题】:Spark SQL - Join that produces an array of matches instead of one row per matchSpark SQL - 生成匹配数组而不是每次匹配一行的连接
【发布时间】:2019-09-12 16:05:24
【问题描述】:

我正在使用 Spark SQL(特别是在 Java 中),在加入时遇到问题,当加入条件有多个匹配项时。

我在每个匹配项的输出中都收到一行,但我希望将它们一起折叠成一个与连接条件匹配的值数组。

假设我有以下两个表:

位置

location  | animal1 | animal2 | animal3
---------------------------------------
australia | badger  | duck    | penguin
thailand  | moose   | penguin | horse
brazil    | zebra   | cow     | pigeon
mexico    | rhino   | donkey  | cat

禁止的动物

banned_animal | banned_animal_ID
--------------------------------
penguin       | 1
zebra         | 2
moose         | 3

我想要做的是组装一个包含位置的表格,然后是一个包含所有禁止在那里的动物 ID 的列。例如,上面的两个表会产生:

location  | banned_animal_IDs
--------------------------------
australia | [1]
thailand  | [1,3]
brazil    | [2]

我不关心数组中 ID 的顺序,如果有多个,那么对于 Thailand 条目,我对 [1,3][3,1] 同样满意

我现在得到的,不是我正在寻找的,是:

location  | banned_animal_IDs
--------------------------------
australia | 1
thailand  | 1
thailand  | 3
brazil    | 2

我这样做的方式:

Dataset<Row> bannedAnimalsByLocation = locations
                .join(bannedAdminals, joinColumn, "INNER");

joinColumn 是被禁止的动物列

locations 表中可能还有很多其他列,所以我不能只在location 列上做一个.groupBy

【问题讨论】:

  • 如果您不想丢失其他列,请尝试使用带有collect_set 的窗口函数作为聚合函数。然后取出条目最多的行并过滤所有其他行。
  • banned_animal 怎么可能是连接列,它没有出现在locations 数据框中?

标签: apache-spark apache-spark-sql


【解决方案1】:

你可以试试这个:

import org.apache.spark.sql.functions._

val locations = sc.parallelize(Seq(
    ("australia", "badger", "duck", "penguin"),
    ("thailand", "moosen", "penguin", "horse"),
    ("brazil", "zebra", "cow", "pigeon"),
    ("mexico", "rhino", "donkey", "cat")
    )).toDF("location","animal1", "animal2", "animal3")


val bannedAdminals = sc.parallelize(Seq(
    ("penguin", "1"),
    ("zebra", "2"),
    ("moosen", "3")
    )).toDF("banned_animal", "banned_animal_ID")



val dfJoined = locations.join(bannedAdminals, locations("animal1") === bannedAdminals("banned_animal")
                                            or locations("animal2") === bannedAdminals("banned_animal")
                                            or locations("animal3") === bannedAdminals("banned_animal"))
                        .select("location", "banned_animal_ID")

dfJoined.groupBy("location").agg(collect_set("banned_animal_ID")).show

结果:

+---------+-----------------------------+
| location|collect_set(banned_animal_ID)|
+---------+-----------------------------+
|australia|                          [1]|
| thailand|                       [3, 1]|
|   brazil|                          [2]|
+---------+-----------------------------+

【讨论】:

    猜你喜欢
    • 2017-09-13
    • 2017-08-25
    • 1970-01-01
    • 2012-02-03
    • 2023-03-13
    • 2020-01-17
    • 2023-03-19
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多