Last active
August 29, 2015 14:07
-
-
Save vdebergue/43e78552bafbb8596890 to your computer and use it in GitHub Desktop.
Linear Equation Solver
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| object Demo extends App { | |
| val formula = Add(Mult(X, Const(2)), Add(Const(-3), X)) | |
| val solutionDicho = Formula.solveForZeroDichotomy(formula) | |
| val solution = Formula.solve0(formula) | |
| println(s"$formula = 0 <=> X = $solution (dichotomy => $solutionDicho)") | |
| val fdiff = Add( Add(Mult(X,X), Mult(X, Const(2))), Const(-2) ) | |
| val extremum = Formula.solve0(fdiff.diff) | |
| println(s"Extremum of $fdiff => X = $extremum") | |
| } | |
| sealed trait Formula { | |
| def eval(x: Double): Double | |
| def diff: Formula | |
| def simplify: Formula | |
| } | |
| sealed trait Commutative { | |
| def left: Formula | |
| def right: Formula | |
| def simplify: Formula = { | |
| if (simplifyCommutative.isDefinedAt((left, right))) { | |
| simplifyCommutative((left, right)) | |
| } else if (simplifyCommutative.isDefinedAt((right, left))) { | |
| simplifyCommutative((right, left)) | |
| } else { | |
| simplifyDefault | |
| } | |
| } | |
| def simplifyDefault: Formula | |
| def simplifyCommutative: PartialFunction[(Formula, Formula), Formula] | |
| } | |
| case object X extends Formula { | |
| def eval(x: Double) = x | |
| def diff = Const(1) | |
| def simplify = this | |
| override def toString(): String = "X" | |
| } | |
| case class Const(value: Double) extends Formula { | |
| def eval(x: Double) = value | |
| def diff = Const(0) | |
| def simplify = if (value == 0) Zero else this | |
| override def toString() = value.toString | |
| } | |
| case object Zero extends Formula { | |
| def eval(x: Double) = 0 | |
| def diff = this | |
| def simplify = this | |
| override def toString(): String = "0" | |
| } | |
| case class Add(left: Formula, right: Formula) extends Formula with Commutative { | |
| def eval(x: Double) = left.eval(x) + right.eval(x) | |
| def diff = Add(left.diff, right.diff) | |
| def simplifyCommutative = { | |
| case (Zero, s) => s | |
| case (Const(c1), Const(c2)) => Const(c1 + c2) | |
| case (X, X) => Mult(Const(2), X) | |
| case (X, Mult(Const(c2), X)) => Mult(X, Const(1 + c2)) | |
| case (Mult(X, Const(c1)), Mult(X, Const(c2))) => Mult(X, Const(c1 + c2)) | |
| case (Add(X, Const(c1)), Const(c2)) => Add(X, Const(c1 + c2)) | |
| case (Add(X, Const(c1)), Mult(X, Const(c2))) => Add(Mult(X, Const(c2 + 1)), Const(c1)) | |
| case (_, X) => Add(X, left.simplify) | |
| } | |
| def simplifyDefault = Add(left.simplify, right.simplify) | |
| override def toString(): String = s"($left + $right)" | |
| } | |
| case class Mult(left: Formula, right: Formula) extends Formula with Commutative { | |
| def eval(x: Double) = left.eval(x) * right.eval(x) | |
| // (u * v)' = u' * v + u * v' | |
| def diff = Add(Mult(left.diff, right), Mult(left, right.diff)) | |
| def simplifyCommutative = { | |
| case (Zero, _) => Zero | |
| case (Const(c1), Const(c2)) => Const(c1 * c2) | |
| case (s, X) => Mult(X, s.simplify) | |
| } | |
| def simplifyDefault = Mult(left.simplify, right.simplify) | |
| override def toString(): String = s"$left * $right" | |
| } | |
| object Formula { | |
| def solveForZeroDichotomy(f: Formula): Option[Double] = { | |
| // choose a and b so that assertion always true | |
| val a = -200 | |
| val b = 200 | |
| val sigma = 1e-3 | |
| assert((f.eval(a) * f.eval(b)) < 0) | |
| def step(a: Double, b: Double, n: Int): Option[Double] = { | |
| if (n > 100) { | |
| None | |
| } else { | |
| val c = (a + b) / 2 | |
| val fc = f.eval(c) | |
| val fa = f.eval(a) | |
| if (fc == 0 || (b - a) < (2 * sigma)) { | |
| Some(c) | |
| } else { | |
| if (fa * fc > 0) step(c, b, n + 1) else step(a, c, n + 1) | |
| } | |
| } | |
| } | |
| step(a, b, 0) | |
| } | |
| def solve0(formula: Formula): Option[Double] = { | |
| def iter(f: Formula, step: Int): Option[Double] = { | |
| println(s"$step: $f") | |
| if (step > 100) | |
| None | |
| else { | |
| f match { | |
| case Const(_) => None | |
| case X => Some(0) | |
| case Add(X, Const(c)) => Some(-c) | |
| case Add(Const(c), X) => Some(-c) | |
| case Mult(X, Const(c)) => Some(0) | |
| case Mult(Const(c), X) => Some(0) | |
| case Add( Mult(X, Const(c1)), Const(c2) ) => Some(-c2 / c1) | |
| case Add( Const(c2), Mult(X, Const(c1)) ) => Some(-c2 / c1) | |
| case _ => iter(f.simplify, step + 1) | |
| } | |
| } | |
| } | |
| iter(formula, 0) | |
| } | |
| } |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment