Fixed the floating parser test

This commit is contained in:
Panupong Pasupat 2017-05-17 18:46:38 -07:00
parent 35aa466d82
commit 79ca2d3758
2 changed files with 38 additions and 0 deletions

View File

@ -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<String, Double> 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

View File

@ -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))");