【问题标题】:How to apply a custom filtering function on a Spark DataFrame如何在 Spark DataFrame 上应用自定义过滤功能
【发布时间】:2017-04-15 10:18:00
【问题描述】:

我有一个如下形式的 DataFrame:

A_DF = |id_A: Int|concatCSV: String|

还有一个:

B_DF = |id_B: Int|triplet: List[String]|

concatCSV 的示例可能如下所示:

"StringD, StringB, StringF, StringE, StringZ"
"StringA, StringB, StringX, StringY, StringZ"
...

triplet 类似于:

("StringA", "StringF", "StringZ")
("StringB", "StringU", "StringR")
...

我想生成A_DFB_DF笛卡尔 集合,例如;

| id_A: Int | concatCSV: String                             | id_B: Int | triplet: List[String]            |
|     14    | "StringD, StringB, StringF, StringE, StringZ" |     21    | ("StringA", "StringF", "StringZ")|
|     14    | "StringD, StringB, StringF, StringE, StringZ" |     45    | ("StringB", "StringU", "StringR")|
|     18    | "StringA, StringB, StringX, StringY, StringG" |     21    | ("StringA", "StringF", "StringZ")|
|     18    | "StringA, StringB, StringX, StringY, StringG" |     45    | ("StringB", "StringU", "StringR")|
|    ...    |                                               |           |                                  |

然后只保留在A_DF("concatCSV") 中出现在B_DF("triplet") 中的至少有两个子字符串(例如StringA, StringB)的记录,即使用filter 排除那些没有的记录满足这个条件

第一个问题是:我可以在不将 DF 转换为 RDD 的情况下执行此操作吗?

第二个问题是:我可以在join 步骤中理想地完成整个事情——作为where 条件吗?

我尝试过类似的实验:

val cartesianRDD = A_DF
   .join(B_DF,"right")
   .where($"triplet".exists($"concatCSV".contains(_)))

where 无法解析。我尝试使用filter 而不是where,但仍然没有运气。此外,由于某些奇怪的原因,cartesianRDD 的类型注释是 SchemaRDD 而不是 DataFrame。我是怎么结束的?最后,我在上面尝试的内容(我写的短代码)是不完整的,因为它只保留来自concatCSV 的一个子字符串的记录,在triplet 中找到。

那么,第三个问题是:我是否应该改用 RDD 并使用自定义过滤功能来解决它​​?

最后,最后一个问题:我可以在 DataFrames 中使用自定义过滤功能吗?

感谢您的帮助。

【问题讨论】:

  • triplet 的类型不应该是List[(String, String, String)]吗?
  • 另外你使用的是什么版本的 Spark?
  • 谢谢,在示例中修复它 - 措辞不好。我使用的是 Spark 1.5.2
  • 事实证明我从 List[String] 开始是正确的,这就是在 scala 中声明类型的方式。
  • 通常,如果你知道一个集合应该有一个固定的长度,那么最好使用一个元组,但这并不重要。

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


【解决方案1】:

函数CROSS JOIN是在Hive中实现的,所以你可以先用Hive SQL做交叉连接:

A_DF.registerTempTable("a")
B_DF.registerTempTable("b")

// sqlContext should be really a HiveContext
val result = sqlContext.sql("SELECT * FROM a CROSS JOIN b") 

然后您可以使用两个udf 过滤到您的预期输出。一个将您的字符串转换为单词数组,另一个将结果数组列的 intersection 和现有列 "triplet"length 提供给我们:

import scala.collection.mutable.WrappedArray
import org.apache.spark.sql.functions.col

val splitArr = udf { (s: String) => s.split(",").map(_.trim) }
val commonLen = udf { (a: WrappedArray[String], 
                       b: WrappedArray[String]) => a.intersect(b).length }

val temp = (result.withColumn("concatArr",
  splitArr(col("concatCSV"))).select(col("*"),
  commonLen(col("triplet"), col("concatArr")).alias("comm"))
  .filter(col("comm") >= 2)
  .drop("comm")
  .drop("concatArr"))

temp.show
+----+--------------------+----+--------------------+
|id_A|           concatCSV|id_B|             triplet|
+----+--------------------+----+--------------------+
|  14|StringD, StringB,...|  21|[StringA, StringF...|
|  18|StringA, StringB,...|  21|[StringA, StringF...|
+----+--------------------+----+--------------------+

【讨论】:

  • 完美答案。谢谢!
猜你喜欢
  • 2013-05-04
  • 1970-01-01
  • 2019-10-15
  • 1970-01-01
  • 1970-01-01
  • 2022-08-14
  • 1970-01-01
  • 1970-01-01
  • 2021-10-30
相关资源
最近更新 更多