【问题标题】:N-Tree Traversal with Scala Causes Stack OverflowScala 的 N-Tree 遍历导致堆栈溢出
【发布时间】:2023-03-26 21:30:02
【问题描述】:

我正在尝试从 N-tree 数据结构返回小部件列表。在我的单元测试中,如果我有大约 2000 个小部件,每个小部件都有一个依赖项,我会遇到堆栈溢出。我认为正在发生的是 for 循环导致我的树遍历不是尾递归的。在scala中写这个更好的方法是什么?这是我的功能:

protected def getWidgetTree(key: String) : ListBuffer[Widget] = {
  def traverseTree(accumulator: ListBuffer[Widget], current: Widget) : ListBuffer[Widget] = {
    accumulator.append(current)

    if (!current.hasDependencies) {
      accumulator
    }  else {
      for (dependencyKey <- current.dependencies) {
        if (accumulator.findIndexOf(_.name == dependencyKey) == -1) {
          traverseTree(accumulator, getWidget(dependencyKey))
        }
      }

      accumulator
    }
  }

  traverseTree(ListBuffer[Widget](), getWidget(key))
}

【问题讨论】:

  • 能否请您将小部件类与测试用例一起发布
  • 这里是 Petro:案例类 Widget(name: String, dependencies: List[String])

标签: scala recursion tree traversal


【解决方案1】:

它不是尾递归的原因是您在函数内部进行了多次递归调用。要尾递归,递归调用只能是函数体中的最后一个表达式。毕竟,重点在于它的工作方式类似于 while 循环(因此,可以转换为循环)。循环不能在一次迭代中多次调用自身。

要进行这样的树遍历,可以使用队列将需要访问的节点结转。

假设我们有这棵树:

//        1
//       / \  
//      2   5
//     / \
//    3   4

用这个简单的数据结构表示:

case class Widget(name: String, dependencies: List[String]) {
  def hasDependencies = dependencies.nonEmpty
}

我们有这张地图指向每个节点:

val getWidget = List(
  Widget("1", List("2", "5")),
  Widget("2", List("3", "4")),
  Widget("3", List()),
  Widget("4", List()),
  Widget("5", List()))
  .map { w => w.name -> w }.toMap

现在我们可以将您的方法重写为尾递归:

def getWidgetTree(key: String): List[Widget] = {
  @tailrec
  def traverseTree(queue: List[String], accumulator: List[Widget]): List[Widget] = {
    queue match {
      case currentKey :: queueTail =>        // the queue is not empty
        val current = getWidget(currentKey)  // get the element at the front
        val newQueueItems =                  // filter out the dependencies already known
          current.dependencies.filterNot(dependencyKey => 
            accumulator.exists(_.name == dependencyKey) && !queue.contains(dependencyKey))
        traverseTree(newQueueItems ::: queueTail, current :: accumulator) // 
      case Nil =>                            // the queue is empty
        accumulator.reverse                  // we're done
    }
  }

  traverseTree(key :: Nil, List[Widget]())
}

并测试一下:

for (k <- 1 to 5)
  println(getWidgetTree(k.toString).map(_.name))

打印:

ListBuffer(1, 2, 3, 4, 5)
ListBuffer(2, 3, 4)
ListBuffer(3)
ListBuffer(4)
ListBuffer(5)

【讨论】:

  • 添加数千个元素时,“accumulator.exists(_.name == dependencyKey)”行可能会使事情变慢一点。我可以做些什么来改进这种查找?
  • @John,将累加器中的所有键保存在缓存中(一个 Set 可以工作)并检查它。这肯定比遍历累加器要好。
  • 我向 traverseTree 添加了一个可变 HashSet 参数作为“keyAccumulator”,我现在检查它而不是 Widget 累加器,这显着提高了性能。如果我索引每个字符串并在整个过程中只使用整数,也许我可以进一步提升它。
【解决方案2】:

对于与@dhg 的答案相同的示例,没有可变状态的等效尾递归函数(ListBuffer)将是:

case class Widget(name: String, dependencies: List[String])

val getWidget = List(
  Widget("1", List("2", "5")),
  Widget("2", List("3", "4")),
  Widget("3", List()),
  Widget("4", List()),
  Widget("5", List())).map { w => w.name -> w }.toMap

def getWidgetTree(key: String): List[Widget] = {
  def addIfNotAlreadyContained(widgetList: List[Widget], widgetNameToAdd: String): List[Widget] = {
    if (widgetList.find(_.name == widgetNameToAdd).isDefined) widgetList
    else                                                      widgetList :+ getWidget(widgetNameToAdd)
  }

  @tailrec
  def traverseTree(currentWidgets: List[Widget], acc: List[Widget]): List[Widget] = currentWidgets match {
    case Nil                                => {
      // If there are no more widgets in this branch return what we've traversed so far
      acc 
    }
    case Widget(name, Nil) :: rest          => {
      // If the first widget is a leaf traverse the rest and add the leaf to the list of traversed
      traverseTree(rest, addIfNotAlreadyContained(acc, name)) 
    }
    case Widget(name, dependencies) :: rest => {
      // If the first widget is a parent, traverse it's children and the rest and add it to the list of traversed
      traverseTree(dependencies.map(getWidget) ++ rest, addIfNotAlreadyContained(acc, name))
    } 
  }

  val root = getWidget(key)
  traverseTree(root.dependencies.map(getWidget) :+ root, List[Widget]())
}

对于同一个测试用例

for (k <- 1 to 5)
  println(getWidgetTree(k.toString).map(_.name).toList.sorted)

给你:

List(2, 3, 4, 5, 1)
List(3, 4, 2)
List(3)
List(4)
List(5)

请注意,这是后序而不是前序遍历。

【讨论】:

    【解决方案3】:

    太棒了!谢谢。我不知道@tailrec 注释。那是一个很酷的小宝石。我不得不稍微调整一下解决方案,因为带有自引用的小部件会导致无限循环。当对 traverseTree 的调用需要一个 List 时,newQueueItems 也是一个 Iterable,所以我必须 toList 那个位。

    def getWidgetTree(key: String): List[Widget] = {
      @tailrec
      def traverseTree(queue: List[String], accumulator: List[Widget]): List[Widget] = {
        queue match {
          case currentKey :: queueTail =>        // the queue is not empty
            val current = getWidget(currentKey)  // get the element at the front
            val newQueueItems =                  // filter out the dependencies already known
              current.dependencies.filter(dependencyKey =>
                !accumulator.exists(_.name == dependencyKey) && !queue.contains(dependencyKey)).toList
            traverseTree(newQueueItems ::: queueTail, current :: accumulator) //
          case Nil =>                            // the queue is empty
            accumulator.reverse                  // we're done
        }
      }
    
      traverseTree(key :: Nil, List[Widget]())
    }
    

    【讨论】:

      猜你喜欢
      • 2012-05-10
      • 2015-07-25
      • 2015-05-21
      • 2014-02-14
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2011-01-13
      相关资源
      最近更新 更多