【问题标题】:how to implement takeUntil with Scala lazy collections如何使用 Scala 惰性集合实现 takeUntil
【发布时间】:2018-08-09 13:13:55
【问题描述】:

我有一个昂贵的函数,我想尽可能少地运行它,满足以下要求:

  • 我有几个输入值要尝试
  • 如果函数返回的值低于给定阈值,我不想尝试其他输入
  • 如果没有结果低于阈值,我想以最小输出取结果

我找不到使用 Iterator 的 takeWhile/dropWhile 的好解决方案,因为我想包含第一个匹配元素。刚刚结束了以下解决方案:

val pseudoResult = Map("a" -> 0.6,"b" -> 0.2, "c" -> 1.0)

def expensiveFunc(s:String) : Double = {
  pseudoResult(s)
}

val inputsToTry = Seq("a","b","c")

val inputIt = inputsToTry.iterator
val results = mutable.ArrayBuffer.empty[(String, Double)]

val earlyAbort = 0.5 // threshold

breakable {
  while (inputIt.hasNext) {
    val name = inputIt.next()
    val res = expensiveFunc(name)
    results += Tuple2(name,res)
    if (res<earlyAbort) break()
  }
}

println(results) // ArrayBuffer((a,0.6), (b,0.2))

val (name, bestResult) = results.minBy(_._2) // (b, 0.2)

如果我设置val earlyAbort = 0.1,结果应该仍然是(b, 0.2),而无需再次评估所有案例。

【问题讨论】:

  • def takeUntil[A](it: Iterator[A])(p: A =&gt; Boolean): Iterator[A] = it.takeWhile(!p(_))
  • @stefanobaghino 我看不出这有什么帮助,你能用我的测试用例分享一个完整的答案吗?

标签: scala lazy-sequences


【解决方案1】:

您可以使用Stream 来实现您正在寻找的东西,记住Stream 是某种惰性集合,它可以按需评估操作。

这是 scala Stream 文档。

你只需要这样做:

val pseudoResult = Map("a" -> 0.6,"b" -> 0.2, "c" -> 1.0)
val earlyAbort = 0.5

def expensiveFunc(s: String): Double = {
  println(s"Evaluating for $s")
  pseudoResult(s)
}

val inputsToTry = Seq("a","b","c")

val results = inputsToTry.toStream.map(input => input -> expensiveFunc(input))
val finalResult = results.find { case (k, res) => res < earlyAbort }.getOrElse(results.minBy(_._2))

如果find 没有得到任何值,您可以使用相同的流来找到最小值,并且函数不会再次评估,这是因为记忆:

Stream 类还采用了记忆化,以便将先前计算的值从 Stream 元素转换为 A 类型的具体值

考虑到如果原始集合为空,则此代码将失败,如果您想支持空集合,则应将 minBy 替换为 sortBy(_._2).headOption 和 getOrElse 替换为 orElse:

val finalResultOpt = results.find { case (k, res) => res < earlyAbort }.orElse(results.sortBy(_._2).headOption)

这个输出是:

评估一个

为 b 评估

finalResult: (String, Double) = (b,0.2)

finalResultOpt: Option[(String, Double)] = Some((b,0.2))

【讨论】:

  • @proximator 你错了,试试看。由于 Stream,它很懒惰。我已经编辑了添加输出的原始帖子,它显然没有评估 c。
  • @proximator 你又错了。说真的,试试看。当 map 中的函数针对流中的元素进行评估时,将不再对其进行评估。请在评论之前尝试我的代码,或者至少尝试了解流的工作原理。
  • 这真的很有趣,我还认为这将评估所有两次没有结果低于 earlyAbort。
  • 不错!这是否意味着它会记住集合的评估部分?
  • 你应该添加来自 ScalaDoc 的引用,上面写着The Stream class also employs memoization such that previously computed values are converted from Stream elements to concrete values of type A
【解决方案2】:

最清晰、最简单的做法是fold 覆盖输入,只传递当前的最佳结果。

val inputIt :Iterator[String] = inputsToTry.iterator
val earlyAbort = 0.5 // threshold

inputIt.foldLeft(("",Double.MaxValue)){ case (low,name) =>
  if (low._2 < earlyAbort) low
  else Seq(low, (name, expensiveFunc(name))).minBy(_._2)
}
//res0: (String, Double) = (b,0.2)

它只调用expensiveFunc() 所需的次数,但它会遍历整个输入迭代器。如果这仍然太繁重(大量输入),那么我会使用尾递归方法。

val inputIt :Iterator[String] = inputsToTry.iterator
val earlyAbort = 0.5 // threshold

def bestMin(low :(String,Double) = ("",Double.MaxValue)) :(String,Double) = {
  if (inputIt.hasNext) {
    val name = inputIt.next()
    val res = expensiveFunc(name)
    if (res < earlyAbort) (name, res)
    else if (res < low._2) bestMin((name,res))
    else bestMin(low)
  } else low
}
bestMin()  //res0: (String, Double) = (b,0.2)

【讨论】:

    【解决方案3】:

    在输入列表中使用视图: 尝试以下方法:

      val pseudoResult = Map("a" -> 0.6, "b" -> 0.2, "c" -> 1.0)
    
      def expensiveFunc(s: String): Double = {
        println(s"executed for ${s}")
        pseudoResult(s)
      }
    
      val inputsToTry = Seq("a", "b", "c")
      val earlyAbort = 0.5 // threshold
    
      def doIt(): List[(String, Double)] = {
    
        inputsToTry.foldLeft(List[(String, Double)]()) {
          case (n, name) =>
    
    
            val res = expensiveFunc(name)
            if(res < earlyAbort) {
              return n++List((name, res))
            }
            n++List((name, res))
        }
    
      }
    
      val (name, bestResult) = doIt().minBy(_._2)
      println(name)
      println(bestResult)
    

    输出:

    executed for a
    executed for b
    b
    0.2
    

    如您所见,只计算 a 和 b,而不计算 c。

    【讨论】:

    • 不,这不起作用,因为我不能保证存在res &lt; earlyAbort,在这种情况下我需要有最小的res,所以我仍然需要一个包含所有结果的外部状态计算
    • 你想在第一次找到结果时中断
    • 是的,但也许我没有结果find将返回None,在这种情况下我需要取最小值的元素。
    • 好的,我现在看到你的问题了。我已经编辑了代码。你能再试一次吗?
    【解决方案4】:

    这是尾递归的用例之一:

      import scala.annotation.tailrec
      val pseudoResult = Map("a" -> 0.6,"b" -> 0.2, "c" -> 1.0)
    
      def expensiveFunc(s:String) : Double = {
        pseudoResult(s)
      }
    
      val inputsToTry = Seq("a","b","c")
    
      val earlyAbort = 0.5 // threshold
    
      @tailrec
      def f(s: Seq[String], result: Map[String, Double] = Map()): Map[String, Double] = s match {
        case Nil => result
        case h::t =>
          val expensiveCalculation = expensiveFunc(h)
          val intermediateResult = result + (h -> expensiveCalculation)
          if(expensiveCalculation < earlyAbort) {
            intermediateResult
          } else {
            f(t, intermediateResult)
          }
      }
      val result = f(inputsToTry)
    
      println(result) // Map(a -> 0.6, b -> 0.2)
    
      val (name, bestResult) = f(inputsToTry).minBy(_._2) // ("b", 0.2)
    

    【讨论】:

    • 虽然这是一个正确的解决方案,但它的可读性不如我原来的解决方案。
    【解决方案5】:

    如果你实现takeUntil 并使用它,如果你没有找到你要找的东西,你仍然需要再次遍历列表来获得最低的一个。可能更好的方法是拥有一个将find 与reduceOption 组合在一起的函数,如果发现某些东西,则提前返回,否则返回将集合减少为单个项目的结果(在您的情况下,找到最小的项目)。

    结果与您使用 Stream 可以达到的效果相当,如之前的回复中强调的那样,但避免了利用记忆化,这对于非常大的集合来说可能很麻烦。

    可能的实现如下:

    import scala.annotation.tailrec
    
    def findOrElse[A](it: Iterator[A])(predicate: A => Boolean,
                                       orElse: (A, A) => A): Option[A] = {
      @tailrec
      def loop(elseValue: Option[A]): Option[A] = {
        if (!it.hasNext) elseValue
        else {
          val next = it.next()
          if (predicate(next)) Some(next)
          else loop(Option(elseValue.fold(next)(orElse(_, next))))
        }
      }
      loop(None)
    }
    

    让我们添加我们的输入来测试这个:

    def f1(in: String): Double = {
      println("calling f1")
      Map("a" -> 0.6, "b" -> 0.2, "c" -> 1.0, "d" -> 0.8)(in)
    }
    
    def f2(in: String): Double = {
      println("calling f2")
      Map("a" -> 0.7, "b" -> 0.6, "c" -> 1.0, "d" -> 0.8)(in)
    }
    
    val inputs = Seq("a", "b", "c", "d")
    

    还有我们的助手:

    def apply[IN, OUT](in: IN, f: IN => OUT): (IN, OUT) =
      in -> f(in)
    
    def threshold[A](a: (A, Double)): Boolean =
      a._2 < 0.5
    
    def compare[A](a: (A, Double), b: (A, Double)): (A, Double) =
      if (a._2 < b._2) a else b
    

    我们现在可以运行它,看看它是怎么回事:

    val r1 = findOrElse(inputs.iterator.map(apply(_, f1)))(threshold, compare)
    val r2 = findOrElse(inputs.iterator.map(apply(_, f2)))(threshold, compare)
    val r3 = findOrElse(Map.empty[String, Double].iterator)(threshold, compare)
    

    r1 是 Some(b, 0.2),r2 是 Some(b, 0.6),r3 是(合理地)None。在第一种情况下,由于我们使用惰性迭代器并提前终止,因此我们只调用了两次f1。

    您可以查看结果并可以使用此代码here on Scastie。

    【讨论】:

      猜你喜欢
      • 2013-02-27
      • 2019-07-08
      • 2017-04-03
      • 2018-01-22
      • 1970-01-01
      • 2011-05-29
      • 1970-01-01
      • 1970-01-01
      • 2011-02-01
      相关资源
      最近更新 更多