【问题标题】:Suggestions to optimize a simple Scala foldLeft over multiple values?建议在多个值上优化简单的 Scala foldLeft?
【发布时间】:2012-02-25 08:29:48
【问题描述】:

我正在从 Java 到 Scala 重新实现一些代码(一个简单的贝叶斯推理算法,但这并不重要)。我想以尽可能高性能的方式实现它,同时通过尽可能避免可变性来保持代码的简洁和功能。

这里是Java代码的sn-p:

    // initialize
    double lP  = Math.log(prior);
    double lPC = Math.log(1-prior);

    // accumulate probabilities from each annotation object into lP and lPC
    for (Annotation annotation : annotations) {
        float prob = annotation.getProbability();
        if (isValidProbability(prob)) {
            lP  += logProb(prob);
            lPC += logProb(1 - prob);
        }
    } 

很简单,对吧?所以我决定第一次尝试使用 Scala foldLeft 和 map 方法。由于我有两个要累加的值,所以累加器是一个元组:

    val initial  = (math.log(prior), math.log(1-prior))
    val probs    = annotations map (_.getProbability)
    val (lP,lPC) = probs.foldLeft(initial) ((r,p) => {
      if(isValidProbability(p)) (r._1 + logProb(p), r._2 + logProb(1-p)) else r
    })

不幸的是,这段代码的执行速度比 Java 慢了大约 5 倍(使用简单且不精确的度量;只是在循环中调用了 10000 次代码)。一个缺陷非常明显;我们遍历列表两次,一次在 map 调用中,另一次在 foldLeft 中。所以这里有一个遍历列表一次的版本。

    val (lP,lPC) = annotations.foldLeft(initial) ((r,annotation) => {
      val  p = annotation.getProbability
      if(isValidProbability(p)) (r._1 + logProb(p), r._2 + logProb(1-p)) else r
    })

这样更好!它的性能比 Java 代码差大约 3 倍。我的下一个预感是,在折叠的每个步骤中创建所有新元组可能都会涉及一些成本。所以我决定尝试一个遍历列表两次但不创建元组的版本。

    val lP = annotations.foldLeft(math.log(prior)) ((r,annotation) => {
       val  p = annotation.getProbability
       if(isValidProbability(p)) r + logProb(p) else r
    })
    val lPC = annotations.foldLeft(math.log(1-prior)) ((r,annotation) => {
      val  p = annotation.getProbability
      if(isValidProbability(p)) r + logProb(1-p) else r
    })

这与以前的版本大致相同(比 Java 版本慢 3 倍)。并不令人惊讶,但我充满希望。

所以我的问题是,有没有更快的方法在 Scala 中实现这个 Java sn-p,同时保持 Scala 代码干净,避免不必要的可变性并遵循 Scala 习语?我确实希望最终在并发环境中使用此代码,因此保持不变性的价值可能超过单线程中较慢的性能。

【问题讨论】:

  • 你在 scala 中有惰性数据结构吗?如果是这样,您应该能够避免多次通过。
  • @Marcin:是的,Scala 集合提供了一个view method,这让这很容易。

标签: scala functional-programming fold optimization


【解决方案1】:

首先,您的部分处罚可能来自您使用的收藏类型。但其中大部分可能是对象创建,您实际上并不能通过运行循环两次来避免,因为数字必须被装箱。

相反,您可以创建一个可变类来为您累积值:

class LogOdds(var lp: Double = 0, var lpc: Double = 0) {
  def *=(p: Double) = {
    if (isValidProbability(p)) {
      lp += logProb(p)
      lpc += logProb(1-p)
    }
    this  // Pass self on so we can fold over the operation
  }
  def toTuple = (lp, lpc)
}

现在虽然您可以不安全地使用它,但您不必这样做。事实上,你可以把它折叠起来。

annotations.foldLeft(new LogOdds()) { (r,ann) => r *= ann.getProbability } toTuple

如果你使用这种模式,所有可变的不安全性都会被隐藏在折叠内;它永远不会逃脱。

现在,您不能进行并行折叠,但您可以进行聚合,这就像折叠需要额外的操作来组合碎片。所以你添加方法

def **(lo: LogOdds) = new LogOdds(lp + lo.lp, lpc + lo.lpc)

LogOdds 然后

annotations.aggregate(new LogOdds())(
  (r,ann) => r *= ann.getProbability,
  (l,r) => l**r
).toTuple

你会很高兴的。

(请随意使用非数学符号,但由于您基本上是在乘以概率,因此乘法符号似乎比合并概率或类似的东西更能直观地了解正在发生的事情。)

【讨论】:

  • 他为什么不能做平行折叠?他只是在添加值,这是可交换的和关联的。
  • @DanielC.Sobral - 因为他需要做一个foldLeft ((U,T)=>U),而不仅仅是弃牌 ((U,U)=>U),而且foldLeft 不能明智地并行累积。这就是aggregate 存在的原因。
  • @Rex - 我也不明白。如果您首先对有效性进行过滤(并忽略 lplpc 的初始化,这只是一个简单的添加),这看起来与我相关。你可以任意并行化什么是Foldable[A : Monoid].sum
  • @oxbow_lakes - 你可以map,然后filter,然后fold,或者你可以aggregate。一步通常比三步快。 aggregate也是并行操作。
  • 啊,好吧,我忘记了map这一步。
【解决方案2】:

您可以实现一个尾递归方法,该方法将由编译器转换为 while 循环,因此应该与 Java 版本一样快。或者,您可以只使用循环 - 如果它只是在方法中使用局部变量(例如,请参阅 Scala 集合源代码中的广泛使用),则没有法律禁止它。

def calc(lst: List[Annotation], lP: Double = 0, lPC: Double = 0): (Double, Double) = {
  if (lst.isEmpty) (lP, lPC)
  else {
    val prob = lst.head.getProbability
    if (isValidProbability(prob)) 
      calc(lst.tail, lP + logProb(prob), lPC + logProb(1 - prob))
    else 
      calc(lst.tail, lP, lPC)
  }
}

折叠的优点是它是可并行化的,这可能会导致它在多核机器上比 Java 版本更快(请参阅其他答案)。

【讨论】:

  • List 不能有效地并行化
  • @oxbow 说得通;如果并行化,最好确保您使用的是具有快速随机访问的类,例如Vector
【解决方案3】:

作为一种旁注:使用view,您可以避免重复遍历列表两次:

val probs = annotations.view.map(_.getProbability).filter(isValidProbability)

val (lP, lPC) = ((logProb(prior), logProb(1 - prior)) /: probs) {
   case ((pa, ca), p) => (pa + logProb(p), ca + logProb(1 - p))
}

这可能不会让您获得比第三个版本更好的性能,但对我来说感觉更优雅。

【讨论】:

  • 感谢 regd view() 的建议,尤其是示例。
  • 在创建视图所需的惰性数据结构时可能会有一些开销。我的简单基准测试表明,其中涉及成本。但我完全同意优雅方面:)
【解决方案4】:

首先,让我们解决性能问题:除了使用 while 循环之外,没有其他方法可以像 Java 一样快速地实现它。基本上,JVM 无法将 Scala 循环优化到优化 Java 循环的程度。其原因甚至是 JVM 人员关心的问题,因为它也妨碍了他们并行库的工作。

现在,回到 Scala 性能,您还可以使用.view 来避免在map 步骤中创建新集合,但我认为map 步骤总是会导致性能变差。问题是,您正在将集合转换为在Double 上参数化的集合,它必须被装箱和拆箱。

但是,有一种可能的优化方法:使其并行化。如果您在annotations 上调用.par 使其成为并行集合,则可以使用fold

val parAnnot = annotations.par
val lP = parAnnot.map(_.getProbability).fold(math.log(prior)) ((r,p) => {
   if(isValidProbability(p)) r + logProb(p) else r
})
val lPC = parAnnot.map(_.getProbability).fold(math.log(1-prior)) ((r,p) => {
  if(isValidProbability(p)) r + logProb(1-p) else r
})

如 Rex 所建议,为避免单独的 map 步骤,请使用 aggregate 而不是 fold

对于奖励积分,您可以使用Future 使两个计算并行运行。不过,我怀疑通过带回元组并一次性运行它会获得更好的性能。您必须对这些东西进行基准测试,看看哪种效果更好。

在并行集合上,它可能会首先使用filter 它来获得有效的注释。或者,也许,collect

val parAnnot = annottions.par.view map (_.getProbability) filter (isValidProbability(_)) force;

val parAnnot = annotations.par collect { case annot if isValidProbability(annot.getProbability) => annot.getProbability }

无论如何,基准测试。

【讨论】:

    【解决方案5】:

    目前无法在没有装箱的情况下与 scala 集合库进行交互。因此,Java 中的原始 doubles 将在 fold 操作中不断被装箱和拆箱,即使您没有将它们包装在 Tuple2 中(专门的 - 但是当然,您已经为每次创建新对象付出了性能开销)。

    【讨论】:

    • 如果你想要性能,最低级别(迭代次数最多的那个)总是要处理原始类型的数组,这真的很烦人。有没有可能的抽象?
    猜你喜欢
    • 1970-01-01
    • 2011-05-09
    • 2018-06-01
    • 2011-03-31
    • 2020-04-14
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2022-07-26
    相关资源
    最近更新 更多