如何优化Scala实现的表达式化简器以完成全量同类项合并?
问题根因分析
- 你的加法表达式是嵌套二叉树结构,当前化简逻辑仅能处理相邻的同类型项,当中间夹入字面量后,后续的同类变量项无法和前序项匹配
- Divide分支存在死递归风险:
case Divide(left, right) => if (left == right) Literal(1.0) else simplify(expr)中else分支传入原expr调用simplify,会无限循环触发栈溢出 - 你当前的递归化简逻辑仅会逐层处理子节点,不会回头重新匹配已经处理过的上层节点,导致后半段的
a+a无法被识别合并
解决思路
要实现完整的同类项合并,需要先把嵌套的加法结构扁平化,再统一聚合同类项,步骤如下:
- 所有Add节点递归展开为平级的项列表,消除嵌套结构
- 对列表中的项分类:
- 字面量统一求和合并为单个Literal
- 变量项按变量名分组,统计出现次数,转换为
系数*变量的结构
- 将聚合后的项重新组装为加法二叉树结构
调整后的代码
首先删掉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
相关产品推荐
相关产品推荐

