diff --git a/build.sbt b/build.sbt index 9099f49..30a9823 100644 --- a/build.sbt +++ b/build.sbt @@ -9,9 +9,6 @@ val z3LibPath = "./lib/native" libraryDependencies ++= Seq( "org.allenai.common" %% "common-core" % "1.4.3", "org.allenai.common" %% "common-testkit" % "1.4.3", - "org.allenai.third_party" % "z3" % "4.4.1-0", - "org.allenai.third_party" % "z3-native-linux" % "4.4.1-0", - "org.allenai.third_party" % "z3-native-macos" % "4.4.1-0", "org.apache.commons" % "commons-compress" % "1.12" ) diff --git a/src/main/scala/org/allenai/iqclid/Z3Solver.scala b/src/main/scala/org/allenai/iqclid/Z3Solver.scala deleted file mode 100644 index 8a2d1cc..0000000 --- a/src/main/scala/org/allenai/iqclid/Z3Solver.scala +++ /dev/null @@ -1,115 +0,0 @@ -package org.allenai.iqclid - -import com.microsoft.z3.ArithExpr -import org.allenai.iqclid.api._ -import org.allenai.iqclid.z3._ - -class Z3Solver { - - def getFunctionTree(depth: Int, nSeq: NumberSequence): Seq[Tree] = { - - var treeIndex = 0 - - ThreadSafeDependencies.withZ3Module { - z3Module => - val smt = new Z3Interface(z3Module, true) - (nSeq.numBaseCases until nSeq.length).foreach { - i => - treeIndex = 0 - smt.add(smt.mkEq(Seq( - functionCSP(depth, nSeq.seq, i), - smt.mkIntConst(nSeq.seq(i))))) - } - - def functionCSP(d: Int, seq: Seq[Int], seqIndex: Int): ArithExpr = { - if (d == 0) { - val isNumber = smt.mkBoolVar(s"isNumber_$treeIndex") - val number = smt.mkIntVar(s"number_$treeIndex") - val isIndex = smt.mkBoolVar(s"isIndex_$treeIndex") - val index = smt.mkIntConst(seqIndex) - val isT1 = smt.mkBoolVar(s"isT1_$treeIndex") - val t1 = smt.mkIntConst(seq(seqIndex - 1)) - val isT2 = smt.mkBoolVar(s"isT2_$treeIndex") - val t2 = smt.mkIntConst(seq(seqIndex - 2)) - val returnVal = smt.mkIntVar(s"r_${seqIndex}_$treeIndex") - // Make sure at least one entity is chosen - smt.add(smt.mkOr(Seq(isNumber, isIndex, isT1, isT2))) - smt.mkImplies(isNumber, smt.mkEq(Seq(number, returnVal))) - smt.mkImplies(isIndex, smt.mkEq(Seq(index, returnVal))) - smt.mkImplies(isT1, smt.mkEq(Seq(t1, returnVal))) - smt.mkImplies(isT2, smt.mkEq(Seq(t2, returnVal))) - treeIndex += 1 - returnVal - } else { - val left = functionCSP(d - 1, seq, seqIndex) - val right = functionCSP(d - 1, seq, seqIndex) - val isPlus = smt.mkBoolVar(s"isPlus_$treeIndex") - val isMinus = smt.mkBoolVar(s"isMinus_$treeIndex") - val isTimes = smt.mkBoolVar(s"isTimes_$treeIndex") - val isDivide = smt.mkBoolVar(s"isDiv_$treeIndex") - // This is used to only use a subtree - val isPickLeft = smt.mkBoolVar(s"isPickLeft_$treeIndex") - val returnVal = smt.mkIntVar(s"r_${seqIndex}_$treeIndex") - smt.add(smt.mkOr(Seq(isPlus, isMinus, isTimes, isDivide, isPickLeft))) - smt.mkImplies(isPlus, smt.mkEq(Seq(smt.mkAdd(Seq(left, right)), returnVal))) - smt.mkImplies(isMinus, smt.mkEq(Seq(smt.mkSub(left, right), returnVal))) - smt.mkImplies(isTimes, smt.mkEq(Seq(smt.mkMul(Seq(left, right)), returnVal))) - smt.mkImplies(isDivide, smt.mkEq(Seq(smt.mkDiv(left, right), returnVal))) - smt.mkImplies(isPickLeft, smt.mkEq(Seq(left, returnVal))) - treeIndex += 1 - returnVal - } - } - - val status = smt.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() - case SmtSatisfiable => - treeIndex = 0 - Seq(modelToAst(smt.extractModel(precision), depth)) - case _ => throw new IllegalStateException("Unrecognized SMT status") - } - } - - def modelToAst(model: Map[String, String], d: Int): Tree = { - if (d == 0) { - val leaf = if (model(s"isNumber_$treeIndex") == "true") { - Number(model(s"number_$treeIndex").toInt) - } else if (model(s"isIndex_$treeIndex") == "true") { - I() - } else if (model(s"isT1_$treeIndex") == "true") { - T(1) - } else if (model(s"isT2_$treeIndex") == "true") { - T(2) - } else { - throw new RuntimeException(s"Something is wrong, no leaf chosen at tree index $treeIndex") - } - treeIndex += 1 - leaf - } else { - val left = modelToAst(model, d - 1) - val right = modelToAst(model, d - 1) - val node = if (model(s"isPlus_$treeIndex") == "true") { - Apply(Plus(), Seq(left, right)) - } else if (model(s"isMinus_$treeIndex") == "true") { - Apply(Minus(), Seq(left, right)) - } else if (model(s"isTimes_$treeIndex") == "true") { - Apply(Times(), Seq(left, right)) - } else if (model(s"isDiv_$treeIndex") == "true") { - Apply(Div(), Seq(left, right)) - } else { - throw new RuntimeException(s"Something is wrong, no op chosen at tree index $treeIndex") - } - treeIndex +=1 - node - } - } - - } - -} diff --git a/src/main/scala/org/allenai/iqclid/api/TreeApi.scala b/src/main/scala/org/allenai/iqclid/api/TreeApi.scala index 15f01be..29548a9 100644 --- a/src/main/scala/org/allenai/iqclid/api/TreeApi.scala +++ b/src/main/scala/org/allenai/iqclid/api/TreeApi.scala @@ -75,9 +75,11 @@ object Evaluator { case (Pow(), Seq(el1, el2)) => Math.pow( evaluateInternal(el1, seqSoFar, index), evaluateInternal(el2, seqSoFar, index)).toInt - case _ => throw new BadTreeException() + case _ => + throw new BadTreeException() } - case _ => throw new BadTreeException() + case _ => + throw new BadTreeException() } } } \ No newline at end of file diff --git a/src/main/scala/org/allenai/iqclid/dataset/IqTest.scala b/src/main/scala/org/allenai/iqclid/dataset/IqTest.scala index 34d2d87..eecc65b 100644 --- a/src/main/scala/org/allenai/iqclid/dataset/IqTest.scala +++ b/src/main/scala/org/allenai/iqclid/dataset/IqTest.scala @@ -1,11 +1,11 @@ package org.allenai.iqclid.dataset -import org.allenai.iqclid.{ DatasetSequence, _ } +import org.allenai.iqclid.{DatasetSequence, _} import org.allenai.iqclid.api._ object IqTest extends Dataset { - val sequences = easy + override lazy val sequences = agiPaper val easy = Seq( DatasetSequence( @@ -62,7 +62,7 @@ object IqTest extends Dataset { val medium = Seq( DatasetSequence( - NumberSequence(Seq(-2, 5, -4, 3, 6), 2), + NumberSequence(Seq(-2, 5, -4, 3, -6), 2), 1, Apply(Minus(), Seq( T(2), Number(2)))), @@ -74,24 +74,30 @@ object IqTest extends Dataset { T(1), Apply(Plus(), Seq(Apply(Times(), Seq( Number(2), I())), Number(1)))))), - DatasetSequence( - NumberSequence(Seq(75, 15, 25, 5, 15), 1), - 3, - Apply(Plus(), - Seq( - Apply(Div(), Seq( - Apply(Times(), Seq( - Apply(Mod(), Seq( - I(), Number(2))), - T(1))), - Number(5))), - Apply(Times(), Seq( - Apply(Mod(), Seq( - Apply(Plus(), Seq( - I(), Number(1))), - Number(2))), - T(1))), - Number(10)))), + + // TODO(row11) fix it later +// DatasetSequence( +// NumberSequence(Seq(75, 15, 25, 5, 15), 1), +// 3, +// Apply(Plus(), +// Seq( +// Apply(Div(), Seq( +// Apply(Times(), +// Seq( +// Apply(Mod(), Seq(I(), Number(2))), +// T(1) +// ) +// ), +// Number(5)) +// ), +// Apply(Times(), +// Seq( +// Apply(Mod(), Seq( +// Apply(Plus(), Seq( +// I(), Number(1))), +// Number(2))), +// Number(10)) +// )))), DatasetSequence( NumberSequence(Seq(1, 2, 6, 24, 120), 1), @@ -143,12 +149,12 @@ object IqTest extends Dataset { DatasetSequence( - NumberSequence(Seq(93, 74, 57, 42, 29), 1), + NumberSequence(Seq(93, 74, 57, 42, 29), 2), 18, Apply(Plus(), Seq( - Apply(Minus(), Seq(T(1), Number(19))), - Apply(Times(), Seq(I(), Number(2))) + Apply(Minus(), Seq(Apply(Times(), Seq(Number(2), T(1))), T(2))), + Number(2) ))), DatasetSequence( @@ -162,10 +168,10 @@ object IqTest extends Dataset { DatasetSequence( NumberSequence(Seq(2, -12, -32, -58, -90), 2), -128, - Apply(Minus(), + Apply(Plus(), Seq( - Apply(Minus(), Seq(T(1), Number(14))), - Apply(Times(), Seq(I(), Number(6))) + Apply(Minus(), Seq(Apply(Times(), Seq(Number(2), T(1))), T(2))), + Number(-6) ))), DatasetSequence( @@ -176,4 +182,161 @@ object IqTest extends Dataset { Number(2) ))) ) + + val eulid = Seq( + DatasetSequence( + NumberSequence(Seq(7, 14, 28, 56, 112), 1), + 224, + Apply(Times(), Seq(T(1), Number(2)))), + + DatasetSequence( + NumberSequence(Seq(-2, -1, 0, 1, 2, 1, 0, -1, -2, -1, 0, 1), 8), + 2, + Apply(Plus(), Seq(T(8), Number(0)))), + + DatasetSequence( + NumberSequence(Seq(1, 2, 3, 1, 2, 3, 1, 2), 0), + 3, + Apply(Plus(), Seq( + Number(1), + Apply( Mod(), Seq(I(), Number(3))) + ))), + + DatasetSequence( + NumberSequence(Seq(1, 2, 3, 1, 2, 3, 1, 2), 0), + 3, + Apply(Plus(), Seq( + Number(1), + Apply( Mod(), Seq(I(), Number(3))) + ))), + + DatasetSequence( + NumberSequence(Seq(-1, 0, 1, -1, 0, 1), 0), + -1, + Apply(Plus(), Seq( + Number(-1), + Apply( Mod(), Seq(I(), Number(3))) + ))), + + DatasetSequence( + NumberSequence(Seq(-1, 1, 2, -1, 1, 2, -1, 1, 2), 3), + -1, + Apply(Plus(), Seq( + Number(0), + T(3) + ))), + + DatasetSequence( + NumberSequence(Seq(1, 2, 1, 2, 1, 2), 0), + 1, + Apply(Plus(), Seq( + Number(1), + Apply( Mod(), Seq(I(), Number(2))) + ))) + ) + + + val agiPaper = Seq( + DatasetSequence( + NumberSequence(Seq(2,5,8,11,14,17,20), 1), + 23, + Apply(Plus(), Seq(T(1), Number(3)))), + + DatasetSequence( + NumberSequence(Seq(25,22,19,16,13,10,7), 1), + 4, + Apply(Plus(), Seq(T(1), Number(-3)))), + + DatasetSequence( + NumberSequence(Seq(8,12,16,20,24,28,32), 1), + 36, + Apply(Plus(), Seq(T(1), Number(4)))), + + DatasetSequence( + NumberSequence(Seq(54,48,42,36,30,24), 1), + 18, + Apply(Plus(), Seq(T(1), Number(-6)))), + + DatasetSequence( + NumberSequence(Seq(28,33,31,36,34,39), 2), + 37, + Apply(Plus(), Seq(T(2), Number(3)))), + + DatasetSequence( + NumberSequence(Seq(6,8,5,7,4,6,3), 2), + 5, + Apply(Plus(), Seq(T(2), Number(-1)))), + + DatasetSequence( + NumberSequence(Seq(9,20,6,17,3,14,0), 2), + 11, + Apply(Plus(), Seq(T(2), Number(-3)))), + + DatasetSequence( + NumberSequence(Seq( 12,15,8,11,4,7,0), 2), + 3, + Apply(Plus(), Seq(T(2), Number(-4)))), + + DatasetSequence( + NumberSequence(Seq(4,11,15,26,41,67), 2), + 108, + Apply(Plus(), Seq(T(2), T(1)))), + + DatasetSequence( + NumberSequence(Seq(3,6,12,24,48,96), 1), + 192, + Apply(Times(), Seq(T(1), Number(2)))), + + DatasetSequence( + NumberSequence(Seq(3,7,15,31,63,127), 1), + 255, + Apply (Plus(), Seq(Apply(Times(), Seq(T(1), Number(2))), Number(1))) + ), + + DatasetSequence( + NumberSequence(Seq(2,3,5,9,17,33,65), 1), + 129, + Apply (Plus(), Seq(Apply(Times(), Seq(T(1), Number(2))), Number(-1))) + ), + + DatasetSequence( + NumberSequence(Seq( 2,12,21,29,36,42,47), 1), + 51, + Apply (Plus(), Seq( + Apply(Plus(), Seq(T(1), Number(11))), + Apply(Times(), Seq(I(), Number(-1))) + ) + ) + ), + + DatasetSequence( + NumberSequence(Seq(148,84,52,36,28,24), 1), + 22, + Apply (Plus(), Seq( + Apply(Div(), Seq(T(1), Number(2))), + Number(10) + ) + ) + ), + + DatasetSequence( + NumberSequence(Seq(148,84,52,36,28,24), 1), + 22, + Apply (Plus(), Seq( + Apply(Div(), Seq(T(1), Number(2))), + Number(10) + ) + ) + ), + + DatasetSequence( + NumberSequence(Seq(2,5,9,19,37,75,149), 1), + 299, + Apply (Plus(), Seq( + Apply(Times(), Seq(T(1), Number(2))), + Apply(Pow(), Seq(Number(-1), Apply(Minus(), Seq(I(), Number(1))))) + ) + ) + ) + ) }