You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优化Scala实现的表达式化简器以完成全量同类项合并?

问题根因分析
  • 你的加法表达式是嵌套二叉树结构,当前化简逻辑仅能处理相邻的同类型项,当中间夹入字面量后,后续的同类变量项无法和前序项匹配
  • Divide分支存在死递归风险:case Divide(left, right) => if (left == right) Literal(1.0) else simplify(expr) 中else分支传入原expr调用simplify,会无限循环触发栈溢出
  • 你当前的递归化简逻辑仅会逐层处理子节点,不会回头重新匹配已经处理过的上层节点,导致后半段的a+a无法被识别合并
解决思路

要实现完整的同类项合并,需要先把嵌套的加法结构扁平化,再统一聚合同类项,步骤如下:

  1. 所有Add节点递归展开为平级的项列表,消除嵌套结构
  2. 对列表中的项分类:
    • 字面量统一求和合并为单个Literal
    • 变量项按变量名分组,统计出现次数,转换为系数*变量的结构
  3. 将聚合后的项重新组装为加法二叉树结构
调整后的代码

首先删掉Expr trait里的simplify方法,把完整的化简逻辑放到Expr伴生对象中,避免调用冲突:

import dsl.Expr.{Add, Divide, Literal, Multiply, Variable, simplifier}

sealed trait Statement

sealed trait Expr extends Statement {
  self =>
  def +(right: Expr): Expr = Expr.Add(self, right)

  def -(right: Expr): Expr = Expr.Add(self, Expr.Negative(right))

  def *(right: Expr): Expr = Expr.Multiply(self, right)

  def /(right: Expr): Expr = Expr.Divide(self, right)

  def evaluate(scope: Map[String, Expr]): Double = self match {
    case expr: Expr =>
      expr match {
        case Expr.Divide(left, right)          => left.evaluate(scope) / right.evaluate(scope)
        case Expr.Multiply(left, right)        => left.evaluate(scope) * right.evaluate(scope)
        case Expr.Add(left, right)             => left.evaluate(scope) + right.evaluate(scope)
        case Expr.Negative(expr)               => -expr.evaluate(scope)
        case Expr.Literal(value)               => value
        case Expr.Variable(name, defaultValue) => scope.getOrElse(name, defaultValue).evaluate(scope)
      }
  }
}

object Expr {
  case class Variable(name: String, defaultValue: Expr) extends Expr
  case class Literal(value: Double)                     extends Expr
  case class Add(left: Expr, right: Expr)               extends Expr
  case class Multiply(left: Expr, right: Expr)          extends Expr
  case class Divide(left: Expr, right: Expr)            extends Expr
  case class Negative(expr: Expr)                       extends Expr
  def apply(expr: Expr): Expr                 = expr

  def simplify(expr: Expr): Expr = expr match {
    // 修复Divide死递归问题
    case Divide(left, right) =>
      val sLeft = simplify(left)
      val sRight = simplify(right)
      if (sLeft == sRight) Literal(1.0) else Divide(sLeft, sRight)
    case Multiply(left, right) => Multiply(simplify(left), simplify(right))
    case Negative(expr) => Negative(simplify(expr))
    // 加法处理核心逻辑
    case Add(_, _) =>
      // 1. 扁平化所有嵌套Add为项列表
      def flattenAdd(e: Expr): List[Expr] = e match {
        case Add(l, r) => flattenAdd(l) ++ flattenAdd(r)
        case other => List(simplify(other))
      }
      val terms = flattenAdd(expr)
      // 2. 分离字面量求和
      val literals = terms.collect { case l: Literal => l.value }
      val literalSum = if (literals.nonEmpty) Some(Literal(literals.sum)) else None
      // 3. 统计变量项的系数
      val varCounts = terms.foldLeft(Map[String, Double]()) { (acc, term) =>
        term match {
          case Variable(name, _) => acc + (name -> (acc.getOrElse(name, 0.0) + 1.0))
          case Multiply(Literal(coef), Variable(name, _)) => acc + (name -> (acc.getOrElse(name, 0.0) + coef))
          case Multiply(Variable(name, _), Literal(coef)) => acc + (name -> (acc.getOrElse(name, 0.0) + coef))
          case _ => acc
        }
      }
      // 4. 把变量计数转为表达式项
      val varTerms = varCounts.map {
        case (name, 1.0) => Variable(name, 1.0): Expr
        case (name, coef) => Multiply(Literal(coef), Variable(name, 1.0)): Expr
      }.toList
      // 5. 合并所有项重新组装Add结构,调整拼接顺序可以控制字面量位置
      val allTerms = varTerms.take(1) ++ literalSum.toList ++ varTerms.drop(1)
      allTerms.reduceLeft(Add)
    case other => other
  }

  def simplifier(expr: Expr): Expr = {
    val simplified = simplify(expr)
    if (simplified == expr) simplified else simplifier(simplified)
  }
  implicit def toExpr(ele: Double): Expr      = Literal(ele)
  implicit def toVariable(name: String): Expr = Variable(name, 1.0)

}

object main extends App {
  import Expr._

  val program: Expr = Variable("a", 1.0) + "a" + 1.0 + "a" + "a"
  val simplified = simplifier(program)
  println(simplified) // 输出:Add(Multiply(Literal(2.0),Variable(a,Literal(1.0))),Add(Literal(1.0),Multiply(Literal(2.0),Variable(a,Literal(1.0))))) 即 2*a + 1 + 2*a
}

如果后续需要支持负数、乘法嵌套加法等更复杂的化简场景,扩展flattenAdd和系数统计的匹配规则即可。

内容的提问来源于stack exchange,提问作者cauchy

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.07 09:12:01