【问题标题】:Spark Dataset equivalent for scala's "collect" taking a partial functionSpark Dataset 等效于 scala 的“收集”,采用部分函数
【发布时间】:2017-06-10 10:22:16
【问题描述】:

常规的 scala 集合有一个漂亮的 collect 方法,它让我可以使用部分函数一次性执行 filter-map 操作。 sparkDatasets 上是否有等效操作?


我喜欢它有两个原因:

  • 语法简洁
  • 它将 filter-map 样式操作减少到单次传递(尽管在 spark 中我猜有一些优化可以为您发现这些东西)

这里有一个例子来说明我的意思。假设我有一个选项序列,我想提取定义的整数并将其加倍(那些在Some 中的整数):

val input = Seq(Some(3), None, Some(-1), None, Some(4), Some(5)) 

方法 1 - collect

input.collect {
  case Some(value) => value * 2
} 
// List(6, -2, 8, 10)

collect 在语法上非常简洁,并且只通过一次。

方法 2 - filter-map

input.filter(_.isDefined).map(_.get * 2)

我可以将这种模式带到 spark 中,因为数据集和数据帧具有类似的方法。

但我不太喜欢这个,因为isDefinedget 对我来说似乎是代码的味道。有一个隐含的假设,即 map 只接收Somes。编译器无法验证这一点。在更大的示例中,开发人员更难发现这种假设,并且开发人员可能会交换过滤器和映射,例如,不会出现语法错误。

方法 3 - fold* 操作

input.foldRight[List[Int]](Nil) {
  case (nextOpt, acc) => nextOpt match {
    case Some(next) => next*2 :: acc
    case None => acc
  }
}

我没有使用足够的 spark 来知道 fold 是否有等价物,所以这可能有点切线。

无论如何,模式匹配、折叠样板和列表的重建都混在一起,很难阅读。


所以总的来说,我发现 collect 语法最好,我希望 spark 有这样的东西。

【问题讨论】:

  • RDDs 和Datasets 上定义的collect 方法用于实现驱动程序中的数据。尽管没有类似于 Collections API collect 方法的东西,但您的直觉是正确的:由于这两个操作都是延迟评估的,因此引擎有机会优化操作并将它们链接起来,以便以最大的局部性执行它们。

标签: scala apache-spark apache-spark-dataset


【解决方案1】:

这里的答案是不正确的,至少在当前的 Spark 中是不正确的。

RDD 实际上有一个 collect 方法,它采用部分函数并将过滤器和映射应用于数据。这与无参数的 .collect() 方法完全不同。请参阅 Spark 源代码 RDD.scala @ line 955:

/**
 * Return an RDD that contains all matching values by applying `f`.
 */
def collect[U: ClassTag](f: PartialFunction[T, U]): RDD[U] = withScope {
  val cleanF = sc.clean(f)
  filter(cleanF.isDefinedAt).map(cleanF)
}

与 RDD.scala @ line 923 中的无参数 .collect() 方法相反,这不会具体化来自 RDD 的数据:

/**
 * Return an array that contains all of the elements in this RDD.
 */
def collect(): Array[T] = withScope {
  val results = sc.runJob(this, (iter: Iterator[T]) => iter.toArray)
  Array.concat(results: _*)
}

在文档中,注意

def collect[U](f: PartialFunction[T, U]): RDD[U]

方法没有有一个关于数据被加载到驱动程序内存中的警告:

https://spark.apache.org/docs/latest/api/scala/index.html#org.apache.spark.rdd.RDD@collect[U](f:PartialFunction[T,U])(implicitevidence$29:scala.reflect.ClassTag[U]):org.apache.spark.rdd.RDD[U]

让这些重载的方法做完全不同的事情对 Spark 来说非常令人困惑。


edit:我的错!我误读了这个问题,我们谈论的是数据集而不是 RDD。尽管如此,接受的答案还是说

”然而,Spark 文档指出,“只有在预期结果数组很小的情况下才应使用此方法,因为所有数据都已加载到驱动程序的内存中。”

这是不正确的!调用 .collect() 的部分函数版本时,数据不会加载到驱动程序的内存中 - 仅在调用无参数版本时。调用 .collect(partial_function) 应该与依次调用 .filter() 和 .map() 具有大致相同的性能,如上面的源代码所示。

【讨论】:

  • 感谢您的回答。问题是关于数据集,而不是 rdd。其他答案之一提到如何将数据集转换为 rdd 然后调用 collect。
  • 对不起,我的错!我将编辑答案,对于某些人来说,了解 .collect() 和 .collect(pf) 之间的区别可能仍然有用。
【解决方案2】:

为了完整起见:

RDD API确实有这样的方法,所以将给定的Dataset / DataFrame转换为RDD,执行collect操作并转换回来总是一个选项,例如:

val dataset = Seq(Some(1), None, Some(2)).toDS()
val dsResult = dataset.rdd.collect { case Some(i) => i * 2 }.toDS()

但是,与在数据集上使用地图和过滤器相比,这可能会表现得更差(原因在@stefanobaghino 的回答中解释)。

对于 DataFrame,这个特定示例(使用 Option)有些误导,因为转换为 DataFrame 实际上会将选项“扁平化”为它们的值(或 null 用于 None),所以等效表达式为:

val dataframe = Seq(Some(1), None, Some(2)).toDF("opt")
dataframe.withColumn("opt", $"opt".multiply(2)).filter(not(isnull($"opt")))

我认为,您对地图操作“假设”其输入的任何内容的担忧较少。

【讨论】:

    【解决方案3】:

    RDDs 和Datasets 上定义的collect 方法用于实现驱动程序中的数据。

    尽管没有类似于 Collections API collect 方法的东西,但您的直觉是正确的:由于这两个操作都是延迟评估的,因此引擎有机会优化操作并将它们链接起来,以便以最大局部性执行它们。

    对于您特别提到的用例,我建议您考虑flatMap,它适用于RDDs 和Datasets:

    // Assumes the usual spark-shell environment
    // sc: SparkContext, spark: SparkSession
    val collection = Seq(Some(1), None, Some(2), None, Some(3))
    val rdd = sc.parallelize(collection)
    val dataset = spark.createDataset(rdd)
    
    // Both operations will yield `Array(2, 4, 6)`
    rdd.flatMap(_.map(_ * 2)).collect
    dataset.flatMap(_.map(_ * 2)).collect
    
    // You can also express the operation in terms of a for-comprehension
    (for (option <- rdd; n <- option) yield n * 2).collect
    (for (option <- dataset; n <- option) yield n * 2).collect
    
    // The same approach is valid for traditional collections as well
    collection.flatMap(_.map(_ * 2))
    for (option <- collection; n <- option) yield n * 2
    

    编辑

    正如在另一个问题中正确指出的那样,RDDs 实际上具有 collect 方法,该方法通过应用部分函数来转换 RDD,就像在普通集合中发生的那样。然而,正如Spark documentation 指出的那样,“只有在预期结果数组很小的情况下才应使用此方法,因为所有数据都已加载到驱动程序的内存中。”

    【讨论】:

    • 感谢@stefanobaghino 的回答!到目前为止,这似乎只剩下我不热衷的方法 2。即使没有收集,是否有更惯用和简洁的方法来解决我的火花数据集示例?
    • 特别是对于您的回答中的情况,flatMap 可以。 :-) val rdd = sc.parallelize(Seq(Some(1), None, Some(2), None, Some(3))); rdd.flatMap(_.map(_ * 2)).collect 将输出 Array(2, 4, 6)。您也可以使用理解。我会将此添加到我的答案中。
    • 感谢您更新您的答案!我忘了for
    • 我认为使用 PartialFunction 的收集没有这个问题...警告在另一种方法中(不带参数收集)。
    【解决方案4】:

    我只是想扩展 stefanobaghino 的答案,包括一个带有案例类的 for 理解示例,因为许多用例可能涉及案例类。

    还有一些选项是 monads,这使得在这种情况下接受的答案非常简单,因为 for 巧妙地丢弃了 None 值,但这种方法不会扩展到像 case 类这样的非 monads:

    case class A(b: Boolean, i: Int, d: Double)
    
    val collection = Seq(A(true, 3), A(false, 10), A(true, -1))
    val rdd = ...
    val dataset = ...
    
    // Select out and double all the 'i' values where 'b' is true:
    for {
      A(b, i, _) <- dataset
      if b
    } yield i * 2
    

    【讨论】:

      【解决方案5】:

      您始终可以创建自己的扩展方法:

      implicit class DatasetOps[T](ds: Dataset[T]) {
      
        def collectt[U](pf: PartialFunction[T, U])(implicit enc: Encoder[U]): Dataset[U] = {
          ds.flatMap(pf.lift(_))
        }
      }
      

      这样:

      // val ds = Dataset(1, 2, 3)
      ds.collectt { case x if x % 2 == 1 => x * 3 }
      // Dataset(3, 9)
      

      请注意,不幸的是,我无法将其命名为 collect(因此是可怕的后缀 t),否则(我认为)签名会与现有的 Dataset#collect 转换 Dataset 的方法发生冲突变成Array

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 2016-01-08
        • 1970-01-01
        • 1970-01-01
        • 2019-06-16
        • 1970-01-01
        • 1970-01-01
        • 2020-09-18
        • 1970-01-01
        相关资源
        最近更新 更多