diff --git a/src/edu/stanford/nlp/sempre/test/ParserTest.java b/src/edu/stanford/nlp/sempre/test/ParserTest.java index 92fa15c..84148e3 100644 --- a/src/edu/stanford/nlp/sempre/test/ParserTest.java +++ b/src/edu/stanford/nlp/sempre/test/ParserTest.java @@ -139,6 +139,33 @@ public class ParserTest { checkRankingArithmetic(new ReinforcementParser(ArithmeticTest().getParserSpec())); } + @Test(groups = "floating") public void checkRankingFloating() { + FloatingParser.opts.defaultIsFloating = true; + FloatingParser.opts.maxDepth = 4; + FloatingParser.opts.useAnchorsOnce = true; + Parser parser = new FloatingParser(new ParseTest(TestUtils.makeArithmeticFloatingGrammar()) { + @Override public void test(Parser parser) {} + }.getParserSpec()); + Params params = new Params(); + Map features = new HashMap<>(); + features.put("rule :: $Operator -> nothing (ConstantFn (lambda y (lambda x (call + (var x) (var y)))))", 1.0); + features.put("rule :: $Operator -> nothing (ConstantFn (lambda y (lambda x (call * (var x) (var y)))))", -1.0); + params.update(features); + /* + * Expected LFs: + * 2 3 + * 2 + 3 3 + 2 + * 2 * 3 3 * 2 + */ + checkNumDerivations(parser, params, "2 and 3", "(number 5)", 6); + + params = new Params(); + features.put("rule :: $Operator -> nothing (ConstantFn (lambda y (lambda x (call + (var x) (var y)))))", -1.0); + features.put("rule :: $Operator -> nothing (ConstantFn (lambda y (lambda x (call * (var x) (var y)))))", 1.0); + params.update(features); + checkNumDerivations(parser, params, "2 and 3", "(number 6)", 6); + } + // TODO(chaganty): verify the parser gradients diff --git a/src/edu/stanford/nlp/sempre/test/TestUtils.java b/src/edu/stanford/nlp/sempre/test/TestUtils.java index 07a4a90..bff60d1 100644 --- a/src/edu/stanford/nlp/sempre/test/TestUtils.java +++ b/src/edu/stanford/nlp/sempre/test/TestUtils.java @@ -33,6 +33,17 @@ public final class TestUtils { return g; } + public static Grammar makeArithmeticFloatingGrammar() { + Grammar g = new Grammar(); + g.addStatement("(rule $Expr ($TOKEN) (NumberFn) (anchored 1))"); + g.addStatement("(rule $Expr ($Expr $Partial) (JoinFn backward))"); + g.addStatement("(rule $Partial ($Operator $Expr) (JoinFn forward))"); + g.addStatement("(rule $Operator (nothing) (ConstantFn (lambda y (lambda x (call + (var x) (var y))))))"); + g.addStatement("(rule $Operator (nothing) (ConstantFn (lambda y (lambda x (call * (var x) (var y))))))"); + g.addStatement("(rule $ROOT ($Expr) (IdentityFn))"); + return g; + } + public static Grammar makeNumberConcatGrammar() { Grammar g = new Grammar(); g.addStatement("(rule $Number ($TOKEN) (NumberFn))");