From f4b84de36b65b39f1665103ec1af3f6e18cef707 Mon Sep 17 00:00:00 2001 From: Ronan Le Bras Date: Thu, 4 Aug 2016 23:11:13 -0700 Subject: [PATCH 1/2] ArithExpr of depth 3; Converting model back to solution --- .../org/allenai/iqclid/z3/SmtSolver.scala | 18 +++ .../org/allenai/iqclid/z3/Z3Interface.scala | 119 ++++++++++-------- .../org/allenai/iqclid/z3/SmtSolverSpec.scala | 33 +++++ 3 files changed, 121 insertions(+), 49 deletions(-) create mode 100644 src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala create mode 100644 src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala diff --git a/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala b/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala new file mode 100644 index 0000000..c320298 --- /dev/null +++ b/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala @@ -0,0 +1,18 @@ +package org.allenai.iqclid.z3 + +import org.allenai.iqclid.NumberSequence +import org.allenai.iqclid.api.{Fitness, Solution, Solver} +import org.allenai.iqclid.z3.ThreadSafeDependencies.Z3Module + +class SmtSolver extends Solver { + override def solve(s: NumberSequence): Seq[Solution]= { + val seq = s.seq + + ThreadSafeDependencies.withZ3Module { + z3Module => + val sol = new Z3Interface(z3Module,true).solveSequence(seq) + println(sol) + sol + } + } +} \ No newline at end of file diff --git a/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala b/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala index bb23d3b..f44cf95 100644 --- a/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala +++ b/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala @@ -2,9 +2,10 @@ package org.allenai.iqclid.z3 import com.microsoft.z3._ import com.microsoft.z3.Symbol +import org.allenai.iqclid.api._ + import scala.collection.mutable import scala.collection.mutable.ArrayBuffer - import org.allenai.iqclid.z3.ThreadSafeDependencies.Z3Module /** various relevant status values after the SMT program is solved */ @@ -362,8 +363,8 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { /** check SMT program for satisfiability */ def check(): SmtStatus = { - println("checking Z3 program for satisfiability") - println(s"SMT program:\n${solver.toString}") +// println("checking Z3 program for satisfiability") +// println(s"SMT program:\n${solver.toString}") // Consider uncommenting the next line. See: // http://stackoverflow.com/questions/15806141/ // keep-getting-unknown-result-with-pattern-usage-in-smtlib-v2-input @@ -382,7 +383,7 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { */ def extractModel(precision: Int): Map[String, String] = { val model = solver.getModel - println("solution extracted: " + model.toString) + //println("solution extracted: " + model.toString) // extract variable assignment in the solution found val solution = boolVarMap.mapValues(model.evaluate(_, false)) ++ intVarMap.mapValues(model.evaluate(_, false)) ++ @@ -407,15 +408,20 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { } } - def solveSequence(s: Seq[Int]) = { + def solveSequence(s: Seq[Int]): Seq[Solution] = { (new SequenceSolver(s)).solve() } class SequenceSolver(s: Seq[Int]) { val seqDomain = ctx.mkFiniteDomainSort("Dom", s.size) - val unaOpDomain = ctx.mkFiniteDomainSort("Dom", 6) + val unaOpDomain = ctx.mkFiniteDomainSort("Dom", 7) val binOpDomain = ctx.mkFiniteDomainSort("Dom", 4) + val level = 3 + val leavesRange = (0 to (scala.math.pow(2,level).toInt-1)) + val lvl1Range = (0 to (scala.math.pow(2,level-1).toInt-1)) + val lvl2Range = (0 to (scala.math.pow(2,level-2).toInt-1)) + def mkSeqEq(fun: FuncDecl): BoolExpr = { val eqTerms = s.indices.map(i => ctx.mkEq(ctx.mkApp(fun, ctx.mkNumeral(i, seqDomain)).asInstanceOf[IntExpr], ctx.mkInt(s(i)))) mkAnd(eqTerms) @@ -438,56 +444,67 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { } def UnaryOp(op: Expr, x: IntExpr): IntExpr = { - ctx.mkITE(ctx.mkEq(op, ctx.mkInt(0)), ctx.mkInt(0), - ctx.mkITE(ctx.mkEq(op, ctx.mkInt(1)), ctx.mkInt(1), - ctx.mkITE(ctx.mkEq(op, ctx.mkInt(2)), ctx.mkInt(2), - ctx.mkITE(ctx.mkEq(op, ctx.mkInt(3)), ctx.mkInt(3), - ctx.mkITE(ctx.mkEq(op, ctx.mkInt(4)), ctx.mkUnaryMinus(x), - x))))).asInstanceOf[IntExpr] + ctx.mkITE(ctx.mkLe(op.asInstanceOf[IntExpr], ctx.mkInt(5)), op.asInstanceOf[IntExpr], + x).asInstanceOf[IntExpr] + } + + def binOpToTree(op: Int, t1: Tree, t2: Tree): Tree = { + op match { + case 0 => + Apply(Plus(),Seq(t1,t2)) + case 1 => + Apply(Times(),Seq(t1,t2)) + case 2 => + Apply(Mod(),Seq(t1,t2)) + case _ => + t2 + } + } + + def unaOpToTree(op: Int): Tree = { + op match { + case n if n <= 5 => + Number(n) + case _ => + I() + } + } + + def formulaToTree(leafOps: Seq[Int], lvl1Ops: Seq[Int], lvl2Ops: Seq[Int], lvl3Op: Int): Tree = { + val leafs = leavesRange.map(i => unaOpToTree(leafOps(i))) + val lvl1Nodes = lvl1Range.map(i => binOpToTree(lvl1Ops(i),leafs(2*i),leafs(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => binOpToTree(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = binOpToTree(lvl3Op, lvl2Nodes(0),lvl2Nodes(1)) + lvl3Node } def mkFormula(x: IntExpr): IntExpr = { - val unaop1 = mkIntVar("unaop1", 0, unaOpDomain.getSize.toInt - 1) - val unaop2 = mkIntVar("unaop2", 0, unaOpDomain.getSize.toInt - 1) - val unaop3 = mkIntVar("unaop3", 0, unaOpDomain.getSize.toInt - 1) - val unaop4 = mkIntVar("unaop4", 0, unaOpDomain.getSize.toInt - 1) - val unaop5 = mkIntVar("unaop5", 0, unaOpDomain.getSize.toInt - 1) - val unaop6 = mkIntVar("unaop6", 0, unaOpDomain.getSize.toInt - 1) - val unaop7 = mkIntVar("unaop7", 0, unaOpDomain.getSize.toInt - 1) - val binop1 = mkIntVar("binop1", 0, binOpDomain.getSize.toInt - 1) - val binop2 = mkIntVar("binop2", 0, binOpDomain.getSize.toInt - 1) - val binop3 = mkIntVar("binop3", 0, binOpDomain.getSize.toInt - 1) - - UnaryOp( - unaop1, - BinaryOp( - binop1, - UnaryOp( - unaop2, - BinaryOp( - binop2, - UnaryOp(unaop4, x), - UnaryOp(unaop5, x) - ) - ), - UnaryOp( - unaop3, - BinaryOp( - binop3, - UnaryOp(unaop6, x), - UnaryOp(unaop7, x) - ) - ) - ) - ) + val leafOps = leavesRange.map(i => mkIntVar(s"leafsOp${i}", 0, unaOpDomain.getSize.toInt - 1)) + val lvl1Ops = lvl1Range.map(i => mkIntVar(s"lvl1Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl2Ops = lvl2Range.map(i => mkIntVar(s"lvl2Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl3Op = mkIntVar(s"lvl3Op0", 0, binOpDomain.getSize.toInt - 1) + + val leaves = leavesRange.map(i => UnaryOp(leafOps(i), x)) + val lvl1Nodes = lvl1Range.map(i => BinaryOp(lvl1Ops(i),leaves(2*i),leaves(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => BinaryOp(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = BinaryOp(lvl3Op,lvl2Nodes(0),lvl2Nodes(1)) + + lvl3Node } // solve 1,2,3,4,5: f(i) = i // solve 1,4,9,16,25: f(i) = i^2 // solve 2,4,6,8,10: f(i) = 2*i // solve 1,2,3,1,2,3: f(i) = (i mod 3) + 1 - def solve() = { + def solve(): Seq[Solution] = { solver.add(mkFormulaEq(mkFormula)) + solver.add(ctx.mkEq(mkIntVar("leafsOp0"),ctx.mkInt(0))) + solver.add(ctx.mkEq(mkIntVar("leafsOp2"),ctx.mkInt(0))) + solver.add(ctx.mkEq(mkIntVar("leafsOp4"),ctx.mkInt(0))) + solver.add(ctx.mkEq(mkIntVar("lvl1Op0"),ctx.mkInt(3))) + solver.add(ctx.mkEq(mkIntVar("lvl1Op1"),ctx.mkInt(3))) + solver.add(ctx.mkEq(mkIntVar("lvl1Op2"),ctx.mkInt(3))) + val status = check() val precision = 16 // number of digits of precision for real values in the model status match { @@ -498,9 +515,13 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { Seq.empty case SmtSatisfiable => val model = extractModel(precision) - println(model) - println("SMT solution: " + model.mkString(", ")) - Seq(model) + val leafOps = leavesRange.map(i => model(s"leafsOp${i}").toInt) + val lvl1Ops = lvl1Range.map(i => model(s"lvl1Op${i}").toInt) + val lvl2Ops = lvl2Range.map(i => model(s"lvl2Op${i}").toInt) + val lvl3Op = model(s"lvl3Op0").toInt + + val tree = formulaToTree(leafOps, lvl1Ops, lvl2Ops, lvl3Op) + Seq(Solution(tree,1)) case _ => throw new IllegalStateException("Unrecognized SMT status") } } diff --git a/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala b/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala new file mode 100644 index 0000000..db9a614 --- /dev/null +++ b/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala @@ -0,0 +1,33 @@ +package org.allenai.iqclid.z3 + +import org.allenai.common.testkit.UnitSpec +import org.allenai.common.testkit.UnitSpec +import org.allenai.iqclid.api._ +import org.allenai.iqclid.{BaselineSearch, NumberSequence} + +class SmtSolverSpec extends UnitSpec{ + + "smt" should "do 1 2 3 4 5" in { + val s = NumberSequence(Seq(1, 2, 3, 4, 5, 6), 1) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 2 4 6 8" in { + val s = NumberSequence(Seq(2,4,6,8), 1) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 1 2 3 1 2 3" in { + val s = NumberSequence(Seq(1,2,3,1,2,3), 1) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 1 4 9 16 25" in { + val s = NumberSequence(Seq(1,4,9,16,25), 1) + val smt = new SmtSolver() + smt.solve(s) + } +} From 703b01a589f166ab6ccaf24cf3f31d6fcf536845 Mon Sep 17 00:00:00 2001 From: Ronan Le Bras Date: Fri, 5 Aug 2016 03:53:39 -0700 Subject: [PATCH 2/2] Recursive relations; Ensemble --- .../org/allenai/iqclid/api/TreeApi.scala | 12 +- .../org/allenai/iqclid/z3/SmtSolver.scala | 30 +- .../org/allenai/iqclid/z3/Z3Interface.scala | 487 +++++++++++++++++- .../org/allenai/iqclid/z3/SmtSolverSpec.scala | 26 +- 4 files changed, 543 insertions(+), 12 deletions(-) diff --git a/src/main/scala/org/allenai/iqclid/api/TreeApi.scala b/src/main/scala/org/allenai/iqclid/api/TreeApi.scala index 15f01be..3cf4f7b 100644 --- a/src/main/scala/org/allenai/iqclid/api/TreeApi.scala +++ b/src/main/scala/org/allenai/iqclid/api/TreeApi.scala @@ -47,6 +47,7 @@ object Evaluator { case I() => index case T(i) => if (index - i < 0) { + println("Wrong index of T(i): " + (index-i)) throw new BadTreeException() } seqSoFar(index - i) @@ -61,6 +62,7 @@ object Evaluator { case (Div(), Seq(el1, el2)) => val denom = evaluateInternal(el2, seqSoFar, index) if (denom == 0) { + println("Unknown operation: " + Apply(op, args)) throw new BadTreeException() } else { evaluateInternal(el1, seqSoFar, index) / evaluateInternal(el2, seqSoFar, index) @@ -68,6 +70,7 @@ object Evaluator { case (Mod(), Seq(el1, el2)) => val denom = evaluateInternal(el2, seqSoFar, index) if (denom == 0) { + println("Unknown operation: " + Apply(op, args)) throw new BadTreeException() } else { evaluateInternal(el1, seqSoFar, index) % evaluateInternal(el2, seqSoFar, index) @@ -75,9 +78,14 @@ object Evaluator { case (Pow(), Seq(el1, el2)) => Math.pow( evaluateInternal(el1, seqSoFar, index), evaluateInternal(el2, seqSoFar, index)).toInt - case _ => throw new BadTreeException() + case _ => { + println("Unknown operation: " + Apply(op, args)) + throw new BadTreeException() + } } - case _ => throw new BadTreeException() + case _ => + println("Unknown tree: " + tree) + throw new BadTreeException() } } } \ No newline at end of file diff --git a/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala b/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala index c320298..a613bc6 100644 --- a/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala +++ b/src/main/scala/org/allenai/iqclid/z3/SmtSolver.scala @@ -1,18 +1,38 @@ package org.allenai.iqclid.z3 import org.allenai.iqclid.NumberSequence -import org.allenai.iqclid.api.{Fitness, Solution, Solver} +import org.allenai.iqclid.api._ import org.allenai.iqclid.z3.ThreadSafeDependencies.Z3Module +import org.allenai.iqclid.api.{Evaluator, Tree} class SmtSolver extends Solver { override def solve(s: NumberSequence): Seq[Solution]= { val seq = s.seq - ThreadSafeDependencies.withZ3Module { - z3Module => - val sol = new Z3Interface(z3Module,true).solveSequence(seq) - println(sol) + val solutionSet = { + val sol = ThreadSafeDependencies.withZ3Module { + z3Module => + new Z3Interface(z3Module, true).solveSequence(seq, 1) + } + if (!sol.isEmpty) { sol + } else { + val sol2 = ThreadSafeDependencies.withZ3Module { + z3Module => + new Z3Interface(z3Module, true).solveSequence(seq, 3) + } + if (!sol2.isEmpty) { + sol2 + } else { + ThreadSafeDependencies.withZ3Module { + z3Module => + new Z3Interface(z3Module, true).solveSequence(seq, 2) + } + } + } } + println(solutionSet) + println(Evaluator.evaluate(solutionSet.head.tree, s.baseCases, s.seq.length + 1)) + solutionSet } } \ No newline at end of file diff --git a/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala b/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala index f44cf95..4d1b8f2 100644 --- a/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala +++ b/src/main/scala/org/allenai/iqclid/z3/Z3Interface.scala @@ -408,8 +408,17 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { } } - def solveSequence(s: Seq[Int]): Seq[Solution] = { - (new SequenceSolver(s)).solve() + def solveSequence(s: Seq[Int], strategy: Int): Seq[Solution] = { + strategy match { + case 1 => + (new SequenceSolver(s)).solve() + case 2 => + (new RecSequenceSolverOrder1(s)).solve() + case 3 => + (new RecSequenceSolverOrder2(s)).solve() + case _ => + (new SequenceSolver(s)).solve() + } } class SequenceSolver(s: Seq[Int]) { @@ -445,7 +454,7 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { def UnaryOp(op: Expr, x: IntExpr): IntExpr = { ctx.mkITE(ctx.mkLe(op.asInstanceOf[IntExpr], ctx.mkInt(5)), op.asInstanceOf[IntExpr], - x).asInstanceOf[IntExpr] + x).asInstanceOf[IntExpr] } def binOpToTree(op: Int, t1: Tree, t2: Tree): Tree = { @@ -509,7 +518,8 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { val precision = 16 // number of digits of precision for real values in the model status match { case SmtUnknown(reason) => - throw new Exception(s"SMT check() returned status UNKNOWN: $reason") + println("Status unknown") + Seq.empty case SmtUnsatisfiable => println("No satisfying assignment found") Seq.empty @@ -526,4 +536,473 @@ class Z3Interface(z3Module: Z3Module, isIntegerProgram: Boolean) { } } } + + class RecSequenceSolver(s: Seq[Int], order: Int) { + val seqDomain = ctx.mkFiniteDomainSort("Dom", s.size) + val unaOpDomain = ctx.mkFiniteDomainSort("Dom", 4) + val binOpDomain = ctx.mkFiniteDomainSort("Dom", 5) + + val level = 3 + val leavesRange = (0 to (scala.math.pow(2,level).toInt-1)) + val lvl1Range = (0 to (scala.math.pow(2,level-1).toInt-1)) + val lvl2Range = (0 to (scala.math.pow(2,level-2).toInt-1)) + + var id = 0 + def getNext(): String = { + id += 1 + s"x${id}" + } + + def mkSeqEq(fun: FuncDecl): BoolExpr = { + val eqTerms = s.indices.map(i => ctx.mkEq(ctx.mkApp(fun, ctx.mkNumeral(i, seqDomain)).asInstanceOf[IntExpr], ctx.mkInt(s(i)))) + mkAnd(eqTerms) + } + + def mkFormulaEq(formula: IntExpr => IntExpr): BoolExpr = { + val eqTerms = s.indices.filter(_>=order).map(i => { + val fi = mkIntVar(s"f_$i") + solver.add(ctx.mkEq(fi, formula(ctx.mkInt(i)))) + ctx.mkEq(fi, ctx.mkInt(s(i))) + }) + mkAnd(eqTerms) + } + + def PrevOp(x: IntExpr): IntExpr = { + val key = s"s_minus_1_${x}" + if( !intVarMap.contains(key) ) { + val Sminus1 = mkIntVar(s"s_minus_1_${x}") + intVarMap.put(key,Sminus1) + val zeroterms = s.indices.map(i => ctx.mkImplies(mkLe(Seq(x, ctx.mkInt(0))), ctx.mkEq(Sminus1, ctx.mkInt(0)))) + solver.add(zeroterms: _*) + val minus1terms = s.indices.filter(_>0).map(i => ctx.mkImplies(mkEq(Seq(x, ctx.mkInt(i))), ctx.mkEq(Sminus1, ctx.mkInt(s(i - 1))))) + solver.add(minus1terms: _*) + } + intVarMap(key) + } + + def SecondPrevOp(x: IntExpr): IntExpr = { + val key = s"s_minus_2_${x}" + if( !intVarMap.contains(key) ) { + val Sminus2 = mkIntVar(s"s_minus_2_${x}") + intVarMap.put(key,Sminus2) + val zeroterms = s.indices.map(i => ctx.mkImplies(mkLe(Seq(x, ctx.mkInt(1))), ctx.mkEq(Sminus2, zeroConst))) + solver.add(zeroterms: _*) + val minus2terms = s.indices.filter(_>1).map(i => ctx.mkImplies(mkEq(Seq(x, ctx.mkInt(i))), ctx.mkEq(Sminus2, ctx.mkInt(s(i - 2))))) + solver.add(minus2terms: _*) + } + intVarMap(key) + } + + def BinaryOp(op: Expr, x: IntExpr, y: IntExpr): IntExpr = { + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(0)), ctx.mkAdd(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(1)), ctx.mkSub(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(2)), ctx.mkMul(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(3)), ctx.mkITE(ctx.mkLe(y, ctx.mkInt(0)), y, ctx.mkMod(x, y)), + y)))).asInstanceOf[IntExpr] + } + + def UnaryOp(opvar: String, op: Expr, x: IntExpr): IntExpr = { + var opvarval = ctx.mkIntConst(opvar) + solver.add(ctx.mkLe(opvarval,ctx.mkInt(9))) + solver.add(ctx.mkGe(opvarval,ctx.mkInt(-1))) + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(0)), ctx.mkInt(0), + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(1)), opvarval, + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(2)), x, + PrevOp(x)))).asInstanceOf[IntExpr] + } + + def binOpToTree(op: Int, t1: Tree, t2: Tree): Tree = { + op match { + case 0 => + Apply(Plus(),Seq(t1,t2)) + case 1 => + Apply(Minus(),Seq(t1,t2)) + case 2 => + Apply(Times(),Seq(t1,t2)) + case 3 => + Apply(Mod(),Seq(t1,t2)) + case _ => + t2 + } + } + + def unaOpToTree(op: Int, opvarval: Int): Tree = { + op match { + case 0 => + Number(0) + case 1 => + Number(opvarval) + case 2 => + I() + case _ => + T(1) + } + } + + def formulaToTree(leafOps: Seq[Int],leafOpsVarVal: Seq[Int], lvl1Ops: Seq[Int], lvl2Ops: Seq[Int], lvl3Op: Int): Tree = { + val leafs = leavesRange.map(i => unaOpToTree(leafOps(i),leafOpsVarVal(i))) + val lvl1Nodes = lvl1Range.map(i => binOpToTree(lvl1Ops(i),leafs(2*i),leafs(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => binOpToTree(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = binOpToTree(lvl3Op, lvl2Nodes(0),lvl2Nodes(1)) + lvl3Node + } + + def mkFormula(x: IntExpr): IntExpr = { + val leafOps = leavesRange.map(i => mkIntVar(s"leafsOp${i}", 0, unaOpDomain.getSize.toInt - 1)) + val lvl1Ops = lvl1Range.map(i => mkIntVar(s"lvl1Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl2Ops = lvl2Range.map(i => mkIntVar(s"lvl2Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl3Op = mkIntVar(s"lvl3Op0", 0, binOpDomain.getSize.toInt - 1) + + val leaves = leavesRange.map(i => UnaryOp(s"leafopvar${i}", leafOps(i), x)) + val lvl1Nodes = lvl1Range.map(i => BinaryOp(lvl1Ops(i),leaves(2*i),leaves(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => BinaryOp(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = BinaryOp(lvl3Op,lvl2Nodes(0),lvl2Nodes(1)) + + lvl3Node + } + + // solve 1,2,3,4,5: f(i) = i + // solve 1,4,9,16,25: f(i) = i^2 + // solve 2,4,6,8,10: f(i) = 2*i + // solve 1,2,3,1,2,3: f(i) = (i mod 3) + 1 + def solve(): Seq[Solution] = { + leavesRange.foreach(i => solver.add(ctx.mkGe(mkIntVar(s"leafopvar${i}"),ctx.mkInt(-5)))) + solver.add(mkFormulaEq(mkFormula)) + solver.add(ctx.mkLe(mkIntVar("leafsOp0"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp1"),ctx.mkInt(3))) + solver.add(ctx.mkLe(mkIntVar("leafsOp2"),ctx.mkInt(1))) + solver.add(ctx.mkLe(mkIntVar("leafsOp3"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp4"),ctx.mkInt(2))) + solver.add(ctx.mkLe(mkIntVar("leafsOp5"),ctx.mkInt(1))) + solver.add(ctx.mkLe(mkIntVar("leafsOp6"),ctx.mkInt(1))) + solver.add(ctx.mkLe(mkIntVar("leafsOp7"),ctx.mkInt(1))) + //(4 to 7).foreach(i => solver.add(ctx.mkLe(mkIntVar(s"leafsOp${i}"),ctx.mkInt(2)))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op0"),ctx.mkInt(4))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op1"),ctx.mkInt(3))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op2"),ctx.mkInt(3))) + + val status = check() + val precision = 16 // number of digits of precision for real values in the model + status match { + case SmtUnknown(reason) => + println("Status unknown") + Seq.empty + case SmtUnsatisfiable => + println("No satisfying assignment found") + Seq.empty + case SmtSatisfiable => + val model = extractModel(precision) + println(model) + val leafOps = leavesRange.map(i => model(s"leafsOp${i}").toInt) + val leafOpsVarVal = leavesRange.map(i => model(s"leafopvar${i}").toInt) + val lvl1Ops = lvl1Range.map(i => model(s"lvl1Op${i}").toInt) + val lvl2Ops = lvl2Range.map(i => model(s"lvl2Op${i}").toInt) + val lvl3Op = model(s"lvl3Op0").toInt + + val tree = formulaToTree(leafOps, leafOpsVarVal, lvl1Ops, lvl2Ops, lvl3Op) + Seq(Solution(tree,1)) + case _ => throw new IllegalStateException("Unrecognized SMT status") + } + } + } + + class RecSequenceSolverOrder1(s: Seq[Int]) { + val a = mkIntVar("a",1,3) + val b = mkIntVar("b",-3,3) + solver.add(mkIsNonZero(b)) + + val seqDomain = ctx.mkFiniteDomainSort("Dom", s.size) + val unaOpDomain = ctx.mkFiniteDomainSort("Dom", 3) + val binOpDomain = ctx.mkFiniteDomainSort("Dom", 5) + + val level = 3 + val leavesRange = (0 to (scala.math.pow(2,level).toInt-1)) + val lvl1Range = (0 to (scala.math.pow(2,level-1).toInt-1)) + val lvl2Range = (0 to (scala.math.pow(2,level-2).toInt-1)) + + var id = 0 + def getNext(): String = { + id += 1 + s"x${id}" + } + + def mkSeqEq(fun: FuncDecl): BoolExpr = { + val eqTerms = s.indices.map(i => ctx.mkEq(ctx.mkApp(fun, ctx.mkNumeral(i, seqDomain)).asInstanceOf[IntExpr], ctx.mkInt(s(i)))) + mkAnd(eqTerms) + } + + def mkFormulaEq(formula: IntExpr => IntExpr): BoolExpr = { + val eqTerms = s.indices.filter(_>=1).map(i => { + val fi = mkIntVar(s"f_$i") + solver.add(ctx.mkEq(fi, formula(ctx.mkInt(i)))) + ctx.mkEq(fi, ctx.mkAdd(ctx.mkMul(a,ctx.mkInt(s(i))),ctx.mkMul(b,ctx.mkInt(s(i-1))))) + }) + mkAnd(eqTerms) + } + + def BinaryOp(op: Expr, x: IntExpr, y: IntExpr): IntExpr = { + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(0)), ctx.mkAdd(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(1)), ctx.mkSub(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(2)), ctx.mkMul(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(3)), ctx.mkITE(ctx.mkLe(y, ctx.mkInt(0)), y, ctx.mkMod(x, y)), + y)))).asInstanceOf[IntExpr] + } + + def UnaryOp(opvar: String, op: Expr, x: IntExpr): IntExpr = { + var opvarval = ctx.mkIntConst(opvar) + solver.add(ctx.mkLe(opvarval,ctx.mkInt(9))) + solver.add(ctx.mkGe(opvarval,ctx.mkInt(-1))) + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(0)), ctx.mkInt(0), + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(1)), opvarval, + x)).asInstanceOf[IntExpr] + } + + def binOpToTree(op: Int, t1: Tree, t2: Tree): Tree = { + op match { + case 0 => + Apply(Plus(),Seq(t1,t2)) + case 1 => + Apply(Minus(),Seq(t1,t2)) + case 2 => + Apply(Times(),Seq(t1,t2)) + case 3 => + Apply(Mod(),Seq(t1,t2)) + case _ => + t2 + } + } + + def unaOpToTree(op: Int, opvarval: Int): Tree = { + op match { + case 0 => + Number(0) + case 1 => + Number(opvarval) + case 2 => + I() + case _ => + T(1) + } + } + + def formulaToTree(leafOps: Seq[Int],leafOpsVarVal: Seq[Int], lvl1Ops: Seq[Int], lvl2Ops: Seq[Int], lvl3Op: Int): Tree = { + val leafs = leavesRange.map(i => unaOpToTree(leafOps(i),leafOpsVarVal(i))) + val lvl1Nodes = lvl1Range.map(i => binOpToTree(lvl1Ops(i),leafs(2*i),leafs(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => binOpToTree(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = binOpToTree(lvl3Op, lvl2Nodes(0),lvl2Nodes(1)) + lvl3Node + } + + def mkFormula(x: IntExpr): IntExpr = { + val leafOps = leavesRange.map(i => mkIntVar(s"leafsOp${i}", 0, unaOpDomain.getSize.toInt - 1)) + val lvl1Ops = lvl1Range.map(i => mkIntVar(s"lvl1Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl2Ops = lvl2Range.map(i => mkIntVar(s"lvl2Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl3Op = mkIntVar(s"lvl3Op0", 0, binOpDomain.getSize.toInt - 1) + + val leaves = leavesRange.map(i => UnaryOp(s"leafopvar${i}", leafOps(i), x)) + val lvl1Nodes = lvl1Range.map(i => BinaryOp(lvl1Ops(i),leaves(2*i),leaves(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => BinaryOp(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = BinaryOp(lvl3Op,lvl2Nodes(0),lvl2Nodes(1)) + + lvl3Node + } + + // solve 1,2,3,4,5: f(i) = i + // solve 1,4,9,16,25: f(i) = i^2 + // solve 2,4,6,8,10: f(i) = 2*i + // solve 1,2,3,1,2,3: f(i) = (i mod 3) + 1 + def solve(): Seq[Solution] = { + leavesRange.foreach(i => solver.add(ctx.mkGe(mkIntVar(s"leafopvar${i}"),ctx.mkInt(-5)))) + solver.add(mkFormulaEq(mkFormula)) + solver.add(ctx.mkLe(mkIntVar("leafsOp0"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp1"),ctx.mkInt(2))) + solver.add(ctx.mkLe(mkIntVar("leafsOp2"),ctx.mkInt(1))) + solver.add(ctx.mkLe(mkIntVar("leafsOp3"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp4"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp5"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp6"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp7"),ctx.mkInt(0))) + //(4 to 7).foreach(i => solver.add(ctx.mkLe(mkIntVar(s"leafsOp${i}"),ctx.mkInt(2)))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op0"),ctx.mkInt(4))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op1"),ctx.mkInt(3))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op2"),ctx.mkInt(3))) + + val status = check() + val precision = 16 // number of digits of precision for real values in the model + status match { + case SmtUnknown(reason) => + throw new Exception(s"SMT check() returned status UNKNOWN: $reason") + case SmtUnsatisfiable => + println("No satisfying assignment found") + Seq.empty + case SmtSatisfiable => + val model = extractModel(precision) + println(model) + val leafOps = leavesRange.map(i => model(s"leafsOp${i}").toInt) + val leafOpsVarVal = leavesRange.map(i => model(s"leafopvar${i}").toInt) + val lvl1Ops = lvl1Range.map(i => model(s"lvl1Op${i}").toInt) + val lvl2Ops = lvl2Range.map(i => model(s"lvl2Op${i}").toInt) + val lvl3Op = model(s"lvl3Op0").toInt + + val aVal = model("a").toInt + val bVal = model("b").toInt + val tree = formulaToTree(leafOps, leafOpsVarVal, lvl1Ops, lvl2Ops, lvl3Op) + val sol = Apply(Div(),Seq(Apply(Minus(),Seq(tree,Apply(Times(),Seq(Number(bVal),T(1))))),Number(aVal))) + Seq(Solution(sol,1)) + case _ => throw new IllegalStateException("Unrecognized SMT status") + } + } + } + + class RecSequenceSolverOrder2(s: Seq[Int]) { + val a = mkIntVar("a",1,3) + val b = mkIntVar("b",-3,3) + val c = mkIntVar("c",-3,3) + solver.add(mkIsNonZero(b)) + solver.add(mkIsNonZero(c)) + + val seqDomain = ctx.mkFiniteDomainSort("Dom", s.size) + val unaOpDomain = ctx.mkFiniteDomainSort("Dom", 3) + val binOpDomain = ctx.mkFiniteDomainSort("Dom", 5) + + val level = 3 + val leavesRange = (0 to (scala.math.pow(2,level).toInt-1)) + val lvl1Range = (0 to (scala.math.pow(2,level-1).toInt-1)) + val lvl2Range = (0 to (scala.math.pow(2,level-2).toInt-1)) + + var id = 0 + def getNext(): String = { + id += 1 + s"x${id}" + } + + def mkSeqEq(fun: FuncDecl): BoolExpr = { + val eqTerms = s.indices.map(i => ctx.mkEq(ctx.mkApp(fun, ctx.mkNumeral(i, seqDomain)).asInstanceOf[IntExpr], ctx.mkInt(s(i)))) + mkAnd(eqTerms) + } + + def mkFormulaEq(formula: IntExpr => IntExpr): BoolExpr = { + val eqTerms = s.indices.filter(_>=2).map(i => { + val fi = mkIntVar(s"f_$i") + solver.add(ctx.mkEq(fi, formula(ctx.mkInt(i)))) + ctx.mkEq(fi, ctx.mkAdd(ctx.mkAdd(ctx.mkMul(a,ctx.mkInt(s(i))),ctx.mkMul(b,ctx.mkInt(s(i-1)))),ctx.mkMul(c,ctx.mkInt(s(i-2))))) + }) + mkAnd(eqTerms) + } + + def BinaryOp(op: Expr, x: IntExpr, y: IntExpr): IntExpr = { + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(0)), ctx.mkAdd(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(1)), ctx.mkSub(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(2)), ctx.mkMul(x, y), + ctx.mkITE(ctx.mkEq(op, ctx.mkInt(3)), ctx.mkITE(ctx.mkLe(y, ctx.mkInt(0)), y, ctx.mkMod(x, y)), + y)))).asInstanceOf[IntExpr] + } + + def UnaryOp(opvar: String, op: Expr, x: IntExpr): IntExpr = { + var opvarval = ctx.mkIntConst(opvar) + solver.add(ctx.mkLe(opvarval,ctx.mkInt(9))) + solver.add(ctx.mkGe(opvarval,ctx.mkInt(-1))) + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(0)), ctx.mkInt(0), + ctx.mkITE(ctx.mkEq(op.asInstanceOf[IntExpr], ctx.mkInt(1)), opvarval, + x)).asInstanceOf[IntExpr] + } + + def binOpToTree(op: Int, t1: Tree, t2: Tree): Tree = { + op match { + case 0 => + Apply(Plus(),Seq(t1,t2)) + case 1 => + Apply(Minus(),Seq(t1,t2)) + case 2 => + Apply(Times(),Seq(t1,t2)) + case 3 => + Apply(Mod(),Seq(t1,t2)) + case _ => + t2 + } + } + + def unaOpToTree(op: Int, opvarval: Int): Tree = { + op match { + case 0 => + Number(0) + case 1 => + Number(opvarval) + case 2 => + I() + case _ => + T(1) + } + } + + def formulaToTree(leafOps: Seq[Int],leafOpsVarVal: Seq[Int], lvl1Ops: Seq[Int], lvl2Ops: Seq[Int], lvl3Op: Int): Tree = { + val leafs = leavesRange.map(i => unaOpToTree(leafOps(i),leafOpsVarVal(i))) + val lvl1Nodes = lvl1Range.map(i => binOpToTree(lvl1Ops(i),leafs(2*i),leafs(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => binOpToTree(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = binOpToTree(lvl3Op, lvl2Nodes(0),lvl2Nodes(1)) + lvl3Node + } + + def mkFormula(x: IntExpr): IntExpr = { + val leafOps = leavesRange.map(i => mkIntVar(s"leafsOp${i}", 0, unaOpDomain.getSize.toInt - 1)) + val lvl1Ops = lvl1Range.map(i => mkIntVar(s"lvl1Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl2Ops = lvl2Range.map(i => mkIntVar(s"lvl2Op${i}", 0, binOpDomain.getSize.toInt - 1)) + val lvl3Op = mkIntVar(s"lvl3Op0", 0, binOpDomain.getSize.toInt - 1) + + val leaves = leavesRange.map(i => UnaryOp(s"leafopvar${i}", leafOps(i), x)) + val lvl1Nodes = lvl1Range.map(i => BinaryOp(lvl1Ops(i),leaves(2*i),leaves(2*i+1))) + val lvl2Nodes = lvl2Range.map(i => BinaryOp(lvl2Ops(i),lvl1Nodes(2*i),lvl1Nodes(2*i+1))) + val lvl3Node = BinaryOp(lvl3Op,lvl2Nodes(0),lvl2Nodes(1)) + + lvl3Node + } + + // solve 1,2,3,4,5: f(i) = i + // solve 1,4,9,16,25: f(i) = i^2 + // solve 2,4,6,8,10: f(i) = 2*i + // solve 1,2,3,1,2,3: f(i) = (i mod 3) + 1 + def solve(): Seq[Solution] = { + leavesRange.foreach(i => solver.add(ctx.mkGe(mkIntVar(s"leafopvar${i}"),ctx.mkInt(-5)))) + solver.add(mkFormulaEq(mkFormula)) + solver.add(ctx.mkLe(mkIntVar("leafsOp0"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp1"),ctx.mkInt(2))) + solver.add(ctx.mkLe(mkIntVar("leafsOp2"),ctx.mkInt(1))) + solver.add(ctx.mkLe(mkIntVar("leafsOp3"),ctx.mkInt(1))) + solver.add(ctx.mkEq(mkIntVar("leafsOp4"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp5"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp6"),ctx.mkInt(0))) + solver.add(ctx.mkLe(mkIntVar("leafsOp7"),ctx.mkInt(0))) + //(4 to 7).foreach(i => solver.add(ctx.mkLe(mkIntVar(s"leafsOp${i}"),ctx.mkInt(2)))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op0"),ctx.mkInt(4))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op1"),ctx.mkInt(3))) + //solver.add(ctx.mkEq(mkIntVar("lvl1Op2"),ctx.mkInt(3))) + + val status = check() + val precision = 16 // number of digits of precision for real values in the model + status match { + case SmtUnknown(reason) => + println("Status unknown") + Seq.empty + case SmtUnsatisfiable => + println("No satisfying assignment found") + Seq.empty + case SmtSatisfiable => + val model = extractModel(precision) + println(model) + val leafOps = leavesRange.map(i => model(s"leafsOp${i}").toInt) + val leafOpsVarVal = leavesRange.map(i => model(s"leafopvar${i}").toInt) + val lvl1Ops = lvl1Range.map(i => model(s"lvl1Op${i}").toInt) + val lvl2Ops = lvl2Range.map(i => model(s"lvl2Op${i}").toInt) + val lvl3Op = model(s"lvl3Op0").toInt + + val aVal = model("a").toInt + val bVal = model("b").toInt + val cVal = model("c").toInt + val tree = formulaToTree(leafOps, leafOpsVarVal, lvl1Ops, lvl2Ops, lvl3Op) + val sol = Apply(Div(),Seq(Apply(Minus(),Seq(Apply(Minus(),Seq(tree,Apply(Times(),Seq(Number(bVal),T(1))))),Apply(Times(),Seq(Number(cVal),T(2))))),Number(aVal))) + Seq(Solution(sol,1)) + case _ => throw new IllegalStateException("Unrecognized SMT status") + } + } + } } diff --git a/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala b/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala index db9a614..384992e 100644 --- a/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala +++ b/src/test/scala/org/allenai/iqclid/z3/SmtSolverSpec.scala @@ -19,8 +19,8 @@ class SmtSolverSpec extends UnitSpec{ smt.solve(s) } + val s = NumberSequence(Seq(1,2,3,1,2,3), 1) "smt" should "do 1 2 3 1 2 3" in { - val s = NumberSequence(Seq(1,2,3,1,2,3), 1) val smt = new SmtSolver() smt.solve(s) } @@ -30,4 +30,28 @@ class SmtSolverSpec extends UnitSpec{ val smt = new SmtSolver() smt.solve(s) } + + "smt" should "do 16, 22, 34, 52, 76" in { + val s = NumberSequence(Seq(16, 22, 34, 52, 76), 2) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 30, 28, 25, 21, 16" in { + val s = NumberSequence(Seq(30, 28, 25, 21, 16), 2) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 123, 135, 148, 160, 173" in { + val s = NumberSequence(Seq(123, 135, 148, 160, 173), 2) + val smt = new SmtSolver() + smt.solve(s) + } + + "smt" should "do 1, 1, 2, 3, 5, 8, 13, 21, 34" in { + val s = NumberSequence(Seq(1, 1, 2, 3, 5, 8, 13, 21, 34), 2) + val smt = new SmtSolver() + smt.solve(s) + } }