diff --git a/README.md b/README.md
index 35b4002..64dba68 100644
--- a/README.md
+++ b/README.md
@@ -1,4 +1,4 @@
-# SEMPRE 2.3.1: Semantic Parsing with Execution
+# SEMPRE 2.4: Semantic Parsing with Execution
## What is semantic parsing?
@@ -62,6 +62,8 @@ SEMPRE has been used in the following papers:
offshoot, and does not use many of the core learning and parsing utiltiies in
SEMPRE. To reproduce those results, check out SEMPRE 1.0.
+Please refer to the [project page](https://nlp.stanford.edu/software/sempre/) for a more complete list.
+
## Where do I go next?
- If you're new to semantic parsing, you can learn more from the [background
@@ -180,3 +182,7 @@ Changes from SEMPRE 2.2 to SEMPRE 2.3:
Changes from SEMPRE 2.3 to SEMPRE 2.3.1:
- Modified the `tables` module to resemble SEMPRE 2.1, effectively making it work again.
+
+Changes from SEMPRE 2.3.1 to SEMPRE 2.4:
+
+- Added the `cprune` package for the paper *Macro Grammars and Holistic Triggering for Efficient Semantic Parsing* (EMNLP 2017).
diff --git a/build.xml b/build.xml
index d033b05..ac72a0a 100644
--- a/build.xml
+++ b/build.xml
@@ -25,7 +25,7 @@
-
+
@@ -78,6 +78,16 @@
+
+
+
+
+
+
+
+
+
+
diff --git a/run b/run
index 2362606..e75ac09 100755
--- a/run
+++ b/run
@@ -634,7 +634,7 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
# Add @cldir=1 to use CodaLab's directory paths
letDefault(:cldir, 0),
# Usual header
- header('core,tables,corenlp'),
+ header('core,tables,corenlp,cprune'),
# Select class
letDefault(:class, 'main'),
sel(:class, {
@@ -649,10 +649,12 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
'alter' => l('edu.stanford.nlp.sempre.tables.alter.BatchTableAlterer', let(:parser, 'serialized')),
'alter-ex' => l('edu.stanford.nlp.sempre.tables.alter.AlteredTablesExecutor', let(:parser, 'serialized')),
'filter' => 'edu.stanford.nlp.sempre.tables.serialize.DumpFilterer',
+ 'column' => 'edu.stanford.nlp.sempre.tables.test.TableColumnAnalyzer',
+ 'execute' => 'edu.stanford.nlp.sempre.tables.test.BatchTableExecutor',
}),
# Fig parameters
selo(:cldir, 'execDir', '_OUTPATH_', '.'),
- o('overwriteExecDir'), o('addToView', 13), o('jarFiles', 'libsempre/*'),
+ o('overwriteExecDir'), o('addToView', 15), o('jarFiles', 'libsempre/*'),
sel(:cldir, l(), '>/dev/null'),
# Set environment for table execution
o('executor', 'tables.lambdadcs.LambdaDCSExecutor'),
@@ -660,6 +662,7 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
o('NumberFn.unitless'), o('NumberFn.alsoTestByConversion'),
o('TypeInference.typeLookup', 'tables.TableTypeLookup'),
o('JoinFn.specializedTypeCheck', false), o('JoinFn.typeInference', true),
+ o('Learner.outputPredValues'),
# Value Evaluator
letDefault(:eval, 'value'),
sel(:eval, {
@@ -672,35 +675,38 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
}),
# Parser
letDefault(:parser, 'floatsize'),
+ o('beamSize', 50),
+ o('useSizeInsteadOfDepth'),
sel(:parser, {
'floatsize' => l(
o('Builder.parser', 'FloatingParser'),
- o('useSizeInsteadOfDepth'),
o('FloatingParser.maxDepth', 15),
nil),
- 'baseline' => l(
- o('Builder.parser', 'tables.baseline.TableBaselineParser'),
- nil),
+ 'baseline' => o('Builder.parser', 'tables.baseline.TableBaselineParser'),
'serialized' => o('Builder.parser', 'tables.serialize.SerializedParser'),
# ACL 2016
'grow-dpd' => l(
o('Builder.parser', 'tables.dpd.DPDParser'),
- o('useSizeInsteadOfDepth'),
o('FloatingParser.maxDepth', 8),
nil),
'grow-float' => l(
o('Builder.parser', 'FloatingParser'),
- o('useSizeInsteadOfDepth'),
o('FloatingParser.maxDepth', 8),
o('FloatingParser.betaReduce'), o('initialFloatingHasZeroDepth'),
nil),
'grow-mix' => l(
o('Builder.parser', 'MixParser'),
o('MixParser.parsers', 'FloatingParser', 'tables.serialize.SerializedParser:train-0xc'),
- o('useSizeInsteadOfDepth'),
o('FloatingParser.maxDepth', 8),
o('FloatingParser.betaReduce'), o('initialFloatingHasZeroDepth'),
nil),
+ # EMNLP 2017
+ 'cprune' => l(
+ o('Builder.parser', 'cprune.CPruneFloatingParser'),
+ o('FloatingParser.maxDepth', 15),
+ o('maxNumNeighbors', 40),
+ o('maxPredictedPatterns', 1000),
+ nil),
}),
o('Parser.verbose', 0),
letDefault(:pruning, 1),
@@ -722,14 +728,14 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
nil),
'editdist-fuzzy' => l(
o('FuzzyMatcher.fuzzyMatcher', 'tables.match.EditDistanceFuzzyMatcher'),
- o('fuzzyMatchSubstring'), o('fuzzyMatchMaxEditDistanceRatio', 0.3),
- o('alsoReturnUnion'), o('alsoMatchPart'),
+ o('fuzzyMatchSubstring'), o('fuzzyMatchMaxEditDistanceRatio', 0.15),
+ o('alsoMatchPart'),
nil),
}),
letDefault(:normalize, 1),
sel(:normalize,
l(),
- l(o('genericDateValue'), o('numberCanStartAnywhere'), o('num2CanStartAnywhere'), o('NumberFn.alsoTestByIsolatedNER')),
+ l(o('genericDateValue'), o('numberCanStartAnywhere'), o('num2CanStartAnywhere')),
nil),
letDefault(:anchor, 1),
sel(:anchor, {
@@ -741,19 +747,24 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
tablesDataPaths,
# Verbosity
o('FeatureVector.ignoreZeroWeight'),
- o('maxPrintedPredictions', 10), o('maxPrintedTrue', 10), o('logFeaturesLimit', 10),
+ o('logFeaturesLimit', 10),
o('LambdaDCSException.noErrorMessage'),
letDefault(:verbose, 0),
sel(:verbose,
- l(),
l(
+ o('maxPrintedPredictions', 1), o('maxPrintedTrue', 1),
+ nil),
+ l(
+ o('maxPrintedPredictions', 10), o('maxPrintedTrue', 10),
o('putCellNameInCanonicalUtterance'), o('showUtterance'),
nil),
l(
+ o('maxPrintedPredictions', 10), o('maxPrintedTrue', 10),
o('putCellNameInCanonicalUtterance'), o('showUtterance'),
o('summarizeRuleTime'), o('summarizeDenotations'),
nil),
l(
+ o('maxPrintedPredictions', 10), o('maxPrintedTrue', 10),
o('putCellNameInCanonicalUtterance'), o('showUtterance'),
o('summarizeRuleTime'), o('summarizeDenotations'),
o('showRules'),
@@ -779,6 +790,7 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
o('combineFromFloatingParser'),
o('maxTrainIters', 3),
o('showValues', false), o('showFirstValue'),
+ o('customExpectedCounts', 'TOP'),
nil),
l(
# for dumping derivations (@class=dump)
@@ -802,10 +814,14 @@ addMode('tables', 'QA on HTML tables', lambda { |e| l(
'some' => l( # Add your own features! (only set up the feature computers)
o('FeatureExtractor.featureComputers', 'tables.features.PhrasePredicateFeatureComputer tables.features.PhraseDenotationFeatureComputer'.split),
nil),
- 'all' => l( # All features
+ 'all' => l( # All ACL 2015 features
o('FeatureExtractor.featureDomains', 'custom-denotation phrase-predicate phrase-denotation headword-denotation missing-predicate'.split),
o('FeatureExtractor.featureComputers', 'tables.features.PhrasePredicateFeatureComputer tables.features.PhraseDenotationFeatureComputer'.split),
nil),
+ 'more' => l( # All ACL 2015 features + more experimental features
+ o('FeatureExtractor.featureDomains', 'custom-denotation phrase-predicate phrase-denotation headword-denotation missing-predicate anchored-entity'.split),
+ o('FeatureExtractor.featureComputers', 'tables.features.PhrasePredicateFeatureComputer tables.features.PhraseDenotationFeatureComputer tables.features.AnchorFeatureComputer'.split),
+ nil),
'baseline' => l( # For the baseline classifier
o('FeatureExtractor.featureDomains', 'custom-denotation phrase-denotation headword-denotation table-baseline'.split),
o('FeatureExtractor.featureComputers', 'tables.baseline.TableBaselineFeatureComputer tables.features.PhraseDenotationFeatureComputer'.split),
@@ -879,6 +895,11 @@ def tablesGrammarPaths
o('Grammar.tags', *'scoped merge-and arithmetic comparison alternative neq yearrange part closedclass-generic scoped-2args-merge-and'.split),
let(:anchor, 2),
nil),
+ # EMNLP 2017
+ 'extended' => l(
+ o('Grammar.inPaths', "#{baseDir}extended.grammar"),
+ o('Grammar.tags', *'alternative movement comparison count aggregate superlative arithmetic merge v-superlative'.split),
+ nil),
}),
nil)
}
@@ -888,13 +909,17 @@ def tablesDataPaths
lambda { |e|
baseDir = ['lib/data/WikiTableQuestions/data/', 'WikiTableQuestions/data/'][e[:cldir]]
csvDir = ['lib/data/WikiTableQuestions/', 'WikiTableQuestions/'][e[:cldir]]
+ nnDir = ['lib/data/nn_0/', 'nn_0/'][e[:cldir]]
datasets = {
'none' => l(),
'train' => o('Dataset.inPaths', "train,#{baseDir}training.examples"),
# Pristine test test
- 'test' => o('Dataset.inPaths',
- "train,#{baseDir}training.examples",
- "test,#{baseDir}pristine-unseen-tables.examples"),
+ 'test' => l(
+ o('Dataset.inPaths',
+ "train,#{baseDir}training.examples",
+ "test,#{baseDir}pristine-unseen-tables.examples"),
+ o('neighborFilePath', "#{nnDir}/exact_nearest_neighbors.all"),
+ nil),
# @data=annotated can be used with @class=check only
'annotated' => o('Dataset.inPaths', "train,#{baseDir}annotated-all.examples"),
'before300' => o('Dataset.inPaths', "train,#{baseDir}training-before300.examples"),
@@ -906,11 +931,13 @@ def tablesDataPaths
"train,#{baseDir}random-split-#{x}-train.examples",
"dev,#{baseDir}random-split-#{x}-dev.examples",
nil),
+ o('neighborFilePath', "#{nnDir}/exact_nearest_neighbors.seed-#{x}.train"),
nil)
end
# That's it!
l(
o('TableKnowledgeGraph.baseCSVDir', csvDir),
+ o('TableValuePreprocessor.taggedFiles', "#{csvDir}/tagged/data/"),
sel(:data, datasets),
nil)
}
@@ -940,10 +967,8 @@ end
############################################################
# {2015-01-18} Generate utterances [Percy]
addMode('genovernight', 'Generate utterances for overnight semantic parsing', lambda { |e| l(
- 'fig/bin/qcreate',
- letDefault(:gen, 0),
- sel(:gen, l()),
- 'java', '-Dmodules=core,overnight', '-Xmx10g', '-cp', 'libsempre/*:lib/*', '-ea', 'edu.stanford.nlp.sempre.overnight.GenerationMain',
+ header('core,overnight'),
+ 'edu.stanford.nlp.sempre.overnight.GenerationMain',
figOpts,
o('JoinFn.typeInference', true),
o('JoinFn.specializedTypeCheck', false),
@@ -959,6 +984,7 @@ addMode('genovernight', 'Generate utterances for overnight semantic parsing', la
o('FeatureExtractor.featureComputers','overnight.OvernightFeatureComputer'),
o('OvernightFeatureComputer.featureDomains', ''),
o('OvernightFeatureComputer.itemAnalysis',false),
+ letDefault(:gen, 1),
sel(:gen,
l( # For debugging the grammar
o('FeatureExtractor.featureDomains', 'denotation'),
diff --git a/src/edu/stanford/nlp/sempre/BooleanValue.java b/src/edu/stanford/nlp/sempre/BooleanValue.java
index ae478b8..497683d 100644
--- a/src/edu/stanford/nlp/sempre/BooleanValue.java
+++ b/src/edu/stanford/nlp/sempre/BooleanValue.java
@@ -19,6 +19,9 @@ public class BooleanValue extends Value {
return tree;
}
+ @Override public String sortString() { return "" + value; }
+ @Override public String pureString() { return "" + value; }
+
@Override public int hashCode() { return Boolean.valueOf(value).hashCode(); }
@Override public boolean equals(Object o) {
if (this == o) return true;
diff --git a/src/edu/stanford/nlp/sempre/DateValue.java b/src/edu/stanford/nlp/sempre/DateValue.java
index 56d186d..3aa71f9 100644
--- a/src/edu/stanford/nlp/sempre/DateValue.java
+++ b/src/edu/stanford/nlp/sempre/DateValue.java
@@ -86,6 +86,7 @@ public class DateValue extends Value {
+ "-" + (month == -1 ? "xx" : String.format("%02d", month))
+ "-" + (day == -1 ? "xx" : String.format("%02d", day));
}
+ @Override public String pureString() { return isoString(); }
@Override public int hashCode() {
int hash = 0x7ed55d16;
diff --git a/src/edu/stanford/nlp/sempre/Example.java b/src/edu/stanford/nlp/sempre/Example.java
index 22f878e..2f93413 100644
--- a/src/edu/stanford/nlp/sempre/Example.java
+++ b/src/edu/stanford/nlp/sempre/Example.java
@@ -198,7 +198,7 @@ public class Example {
public void preprocess() {
this.languageInfo = LanguageAnalyzer.getSingleton().analyze(this.utterance);
- this.targetValue = TargetValuePreprocessor.getSingleton().preprocess(this.targetValue);
+ this.targetValue = TargetValuePreprocessor.getSingleton().preprocess(this.targetValue, this);
}
public void log() {
diff --git a/src/edu/stanford/nlp/sempre/ExampleUtils.java b/src/edu/stanford/nlp/sempre/ExampleUtils.java
index 9bc16dd..91a317b 100644
--- a/src/edu/stanford/nlp/sempre/ExampleUtils.java
+++ b/src/edu/stanford/nlp/sempre/ExampleUtils.java
@@ -119,6 +119,29 @@ public final class ExampleUtils {
out.close();
}
+ public static void writePredictionTSV(int iter, String group, Example ex) {
+ String basePath = "preds-iter" + iter + "-" + group + ".tsv";
+ String outPath = Execution.getFile(basePath);
+ if (outPath == null) return;
+ PrintWriter out = IOUtils.openOutAppendHard(outPath);
+
+ List fields = new ArrayList<>();
+ fields.add(ex.id);
+
+ if (!ex.predDerivations.isEmpty()) {
+ Derivation deriv = ex.predDerivations.get(0);
+ if (deriv.value instanceof ListValue) {
+ List values = ((ListValue) deriv.value).values;
+ for (Value v : values) {
+ fields.add(v.pureString().replaceAll("\\s+", " ").trim());
+ }
+ }
+ }
+
+ out.println(String.join("\t", fields));
+ out.close();
+ }
+
//read lisptree and write json
public static void main(String[] args) {
Dataset dataset = new Dataset();
diff --git a/src/edu/stanford/nlp/sempre/FloatingParser.java b/src/edu/stanford/nlp/sempre/FloatingParser.java
index 4f433f2..6f7e5eb 100644
--- a/src/edu/stanford/nlp/sempre/FloatingParser.java
+++ b/src/edu/stanford/nlp/sempre/FloatingParser.java
@@ -84,18 +84,34 @@ public class FloatingParser extends Parser {
@Option(gloss = "DEBUG: Print amount of time spent on each rule")
public boolean summarizeRuleTime = false;
@Option(gloss = "Stop the parser if it has used more than this amount of time (in seconds)")
- public int maxFloatingParsingTime = 600;
+ public int maxFloatingParsingTime = Integer.MAX_VALUE;
}
public static Options opts = new Options();
- protected List orderedFloatingRules;
- public List getOrderedFloatingRules() { return orderedFloatingRules; }
+ public boolean earlyStopOnConsistent = false;
+ public int earlyStopOnNumDerivs = -1;
public FloatingParser(Spec spec) {
super(spec);
}
+ /**
+ * Set early stopping criteria
+ *
+ * @param onConsistent
+ * Stop when a consistent derivation is found. (Only triggered when computeExpectedCounts = true)
+ * @param onNumDerivs
+ * Stop when the number of featurized derivations exceed this number (set to -1 to disable)
+ * @return
+ * this
+ */
+ public FloatingParser setEarlyStopping(boolean onConsistent, int onNumDerivs) {
+ this.earlyStopOnConsistent = onConsistent;
+ this.earlyStopOnNumDerivs = onNumDerivs;
+ return this;
+ }
+
/**
* computeCatUnaryRules, but do not topologically sort floating rules
*/
@@ -379,39 +395,54 @@ class FloatingParserState extends ParserState {
StopWatch stopWatch = new StopWatch().start();
String rhs1 = rule.rhs.get(0);
String rhs2 = rule.rhs.get(1);
- if (!Rule.isCat(rhs1) || !Rule.isCat(rhs2))
- throw new RuntimeException("Floating rules with > 1 arguments cannot have tokens on the RHS: " + rule);
- if (FloatingParser.opts.useSizeInsteadOfDepth) {
- derivLoop:
- for (int depth1 = 0; depth1 < depth; depth1++) { // sizes must add up to depth-1 (actually size-1)
- int depth2 = depth - 1 - depth1;
- for (ChildDerivationsGroup group : getFilteredDerivations(rule, floatingCell(rhs1, depth1), floatingCell(rhs2, depth2)))
- for (Derivation deriv1 : group.derivations1)
- for (Derivation deriv2 : group.derivations2)
- if (!applyFloatingRule(rule, depth, deriv1, deriv2, deriv1.canonicalUtterance + " " + deriv2.canonicalUtterance))
- break derivLoop;
- }
- } else {
- {
- derivLoop:
- for (int subDepth = 0; subDepth < depth; subDepth++) { // depth-1 <=depth-1
- for (ChildDerivationsGroup group : getFilteredDerivations(rule, floatingCell(rhs1, depth - 1), floatingCell(rhs2, subDepth)))
- for (Derivation deriv1 : group.derivations1)
- for (Derivation deriv2 : group.derivations2)
- if (!applyFloatingRule(rule, depth, deriv1, deriv2, deriv1.canonicalUtterance + " " + deriv2.canonicalUtterance))
- break derivLoop;
- }
+ if (!Rule.isCat(rhs1) && !Rule.isCat(rhs2)) { // token token
+ if (depth == (FloatingParser.opts.initialFloatingHasZeroDepth ? 0 : 1)) {
+ applyFloatingRule(rule, depth, null, null, rhs1 + " " + rhs2);
}
- {
+
+ } else if (!Rule.isCat(rhs1) && Rule.isCat(rhs2)) { // token $Cat
+ List derivations = getDerivations(floatingCell(rhs2, depth - 1));
+ for (Derivation deriv : derivations)
+ applyFloatingRule(rule, depth, deriv, null, rhs1 + " " + deriv.canonicalUtterance);
+
+ } else if (Rule.isCat(rhs1) && !Rule.isCat(rhs2)) { // $Cat token
+ List derivations = getDerivations(floatingCell(rhs1, depth - 1));
+ for (Derivation deriv : derivations)
+ applyFloatingRule(rule, depth, deriv, null, deriv.canonicalUtterance + " " + rhs2);
+
+ } else { // $Cat $Cat
+ if (FloatingParser.opts.useSizeInsteadOfDepth) {
derivLoop:
- for (int subDepth = 0; subDepth < depth - 1; subDepth++) { // categories = new HashSet<>();
+ for (Rule rule : parser.grammar.rules)
+ categories.add(rule.lhs);
+
+ if (Parser.opts.verbose >= 1)
+ LogInfo.begin_track_printAll("Anchored");
+ // Build up anchored derivations (like the BeamParser)
+ int numTokens = ex.numTokens();
+ for (int len = 1; len <= numTokens; len++) {
+ for (int i = 0; i + len <= numTokens; i++) {
+ buildAnchored(i, i + len);
+ for (String cat : categories) {
+ String cell = anchoredCell(cat, i, i + len).toString();
+ pruneCell(cell, chart.get(cell));
+ }
+ }
+ }
+ if (Parser.opts.verbose >= 1)
+ LogInfo.end_track();
+
+ // Build up floating derivations
+ for (int depth = (FloatingParser.opts.initialFloatingHasZeroDepth ? 0 : 1); depth <= FloatingParser.opts.maxDepth; depth++) {
+ if (Parser.opts.verbose >= 1)
+ LogInfo.begin_track_printAll("%s = %d", FloatingParser.opts.useSizeInsteadOfDepth ? "SIZE" : "DEPTH", depth);
+ buildFloating(depth);
+ for (String cat : categories) {
+ String cell = floatingCell(cat, depth).toString();
+ pruneCell(cell, chart.get(cell));
+ }
+ if (Parser.opts.verbose >= 1)
+ LogInfo.end_track();
+ // Early stopping
+ if (computeExpectedCounts && ((FloatingParser) parser).earlyStopOnConsistent) {
+ // Consistent derivation found?
+ String cell = floatingCell(Rule.rootCat, depth).toString();
+ List rootDerivs = chart.get(cell);
+ if (rootDerivs != null) {
+ for (Derivation rootDeriv : rootDerivs) {
+ rootDeriv.ensureExecuted(parser.executor, ex.context);
+ if (parser.valueEvaluator.getCompatibility(ex.targetValue, rootDeriv.value) == 1) {
+ LogInfo.logs("Early stopped: consistent derivation found at depth = %d", depth);
+ return;
+ }
+ }
+ }
+ }
+ if (((FloatingParser) parser).earlyStopOnNumDerivs > 0) {
+ // Too many derivations generated?
+ if (numOfFeaturizedDerivs > ((FloatingParser) parser).earlyStopOnNumDerivs) {
+ LogInfo.logs("Early stopped: number of derivations exceeded at depth = %d", depth);
+ return;
+ }
+ }
+ }
+ }
+ }
+
+ public void buildDerivations() {
+ DerivationBuilder derivBuilder = new DerivationBuilder();
+ if (FloatingParser.opts.maxFloatingParsingTime == Integer.MAX_VALUE) {
+ derivBuilder.run();
+ } else {
+ Thread parsingThread = new Thread(derivBuilder);
+ parsingThread.start();
+ try {
+ parsingThread.join(FloatingParser.opts.maxFloatingParsingTime * 1000);
+ if (parsingThread.isAlive()) {
+ // This will only interrupt first or second passes, not the final candidate collection.
+ LogInfo.warnings("Parsing time exceeded %d seconds. Will now interrupt ...", FloatingParser.opts.maxFloatingParsingTime);
+ timeout = true;
+ parsingThread.interrupt();
+ parsingThread.join();
+ }
+ } catch (InterruptedException e) {
+ e.printStackTrace();
+ LogInfo.fails("FloatingParser error: %s", e);
+ }
+ }
+ evaluation.add("timeout", timeout);
+ }
+
+ // ============================================================
+ // Main entry point
+ // ============================================================
+
@Override public void infer() {
LogInfo.begin_track_printAll("FloatingParser.infer()");
ruleTime = new HashMap<>();
- // Base case ($TOKEN, $PHRASE)
- for (Derivation deriv : gatherTokenAndPhraseDerivations()) {
- addToChart(anchoredCell(deriv.cat, deriv.start, deriv.end), deriv);
- addToChart(floatingCell(deriv.cat, 0), deriv);
- }
-
- Set categories = new HashSet<>();
- for (Rule rule : parser.grammar.rules)
- categories.add(rule.lhs);
-
- if (Parser.opts.verbose >= 1)
- LogInfo.begin_track_printAll("Anchored");
- // Build up anchored derivations (like the BeamParser)
- int numTokens = ex.numTokens();
- for (int len = 1; len <= numTokens; len++) {
- for (int i = 0; i + len <= numTokens; i++) {
- buildAnchored(i, i + len);
- for (String cat : categories) {
- String cell = anchoredCell(cat, i, i + len).toString();
- pruneCell(cell, chart.get(cell));
- }
- }
- }
- if (Parser.opts.verbose >= 1)
- LogInfo.end_track();
-
- // Build up floating derivations
- // Timeout if taking too long
- timeout = false;
- Thread parsingThread = new Thread(new Runnable() {
- @Override
- public void run() {
- for (int depth = (FloatingParser.opts.initialFloatingHasZeroDepth ? 0 : 1); depth <= FloatingParser.opts.maxDepth; depth++) {
- if (Parser.opts.verbose >= 1)
- LogInfo.begin_track_printAll("%s = %d", FloatingParser.opts.useSizeInsteadOfDepth ? "SIZE" : "DEPTH", depth);
- buildFloating(depth);
- for (String cat : categories) {
- String cell = floatingCell(cat, depth).toString();
- pruneCell(cell, chart.get(cell));
- }
- if (Parser.opts.verbose >= 1)
- LogInfo.end_track();
- }
- }
- });
- parsingThread.start();
- try {
- parsingThread.join(FloatingParser.opts.maxFloatingParsingTime * 1000);
- if (parsingThread.isAlive()) {
- // This will only interrupt first or second passes, not the final candidate collection.
- LogInfo.warnings("Parsing time exceeded %d seconds. Will now interrupt ...", FloatingParser.opts.maxFloatingParsingTime);
- timeout = true;
- parsingThread.interrupt();
- parsingThread.join();
- }
- } catch (InterruptedException e) {
- e.printStackTrace();
- LogInfo.fails("DPParser error: %s", e);
- }
- evaluation.add("timeout", timeout);
+ buildDerivations();
if (FloatingParser.opts.summarizeRuleTime) summarizeRuleTime();
diff --git a/src/edu/stanford/nlp/sempre/FuzzyMatchFn.java b/src/edu/stanford/nlp/sempre/FuzzyMatchFn.java
index 6c98126..9f427c5 100644
--- a/src/edu/stanford/nlp/sempre/FuzzyMatchFn.java
+++ b/src/edu/stanford/nlp/sempre/FuzzyMatchFn.java
@@ -39,6 +39,9 @@ public class FuzzyMatchFn extends SemanticFn {
}
}
+ public FuzzyMatchFnMode getMode() { return mode; }
+ public boolean getMatchAny() { return matchAny; }
+
@Override
public DerivationStream call(Example ex, Callable c) {
return new LazyFuzzyMatchFnDerivs(ex, c, mode, matchAny);
diff --git a/src/edu/stanford/nlp/sempre/Grammar.java b/src/edu/stanford/nlp/sempre/Grammar.java
index bb94b36..72cf0c8 100644
--- a/src/edu/stanford/nlp/sempre/Grammar.java
+++ b/src/edu/stanford/nlp/sempre/Grammar.java
@@ -44,7 +44,7 @@ public class Grammar {
// All the rules in the grammar. Each parser can read these and transform
// them however the parser wishes.
// This contains binarized rules
- ArrayList rules = new ArrayList<>();
+ protected ArrayList rules = new ArrayList<>();
public List getRules() { return rules; }
Map macros = new HashMap<>(); // Map from macro name to its replacement value
@@ -260,7 +260,7 @@ public class Grammar {
return cat;
}
- private void interpretRule(LispTree tree) {
+ protected void interpretRule(LispTree tree) {
if (tree.children.size() < 4)
throw new RuntimeException("Invalid rule: " + tree);
@@ -360,7 +360,7 @@ public class Grammar {
// Generate intermediate categories for binarization.
public static final String INTERMEDIATE_PREFIX = "$Intermediate";
- private int freshCatIndex = 0;
+ protected int freshCatIndex = 0;
private String generateFreshCat() {
freshCatIndex++;
return INTERMEDIATE_PREFIX + freshCatIndex;
@@ -368,6 +368,9 @@ public class Grammar {
public static boolean isIntermediate(String cat) {
return cat.startsWith(INTERMEDIATE_PREFIX);
}
+ public int getFreshCatIndex() {
+ return freshCatIndex;
+ }
// Create multiple versions of this rule if there are optional RHS.
// Restriction: must be able to split the RHS into two halves, each of
diff --git a/src/edu/stanford/nlp/sempre/Learner.java b/src/edu/stanford/nlp/sempre/Learner.java
index 6b2ce83..5dd0f3f 100644
--- a/src/edu/stanford/nlp/sempre/Learner.java
+++ b/src/edu/stanford/nlp/sempre/Learner.java
@@ -27,6 +27,8 @@ public class Learner {
@Option(gloss = "Write predDerivations to examples file (huge)")
public boolean outputPredDerivations = false;
+ @Option(gloss = "Write predicted values to a TSV file")
+ public boolean outputPredValues = false;
@Option(gloss = "Dump all features and compatibility scores")
public boolean dumpFeaturesAndCompatibility = false;
@@ -148,7 +150,7 @@ public class Learner {
params.update(counts);
LogInfo.end_track();
}
-
+
public void onlineLearnExampleByFormula(Example ex, List formulas) {
HashMap counts = new HashMap<>();
for (Derivation deriv : ex.predDerivations)
@@ -222,8 +224,11 @@ public class Learner {
addFeedback(ex);
// Write out examples and predictions
- if (opts.outputPredDerivations && Builder.opts.parser.equals("FloatingParser")) {
- ExampleUtils.writeParaphraseSDF(iter, group, ex, opts.outputPredDerivations);
+ if (opts.outputPredDerivations) {
+ ExampleUtils.writeParaphraseSDF(iter, group, ex, true);
+ }
+ if (opts.outputPredValues) {
+ ExampleUtils.writePredictionTSV(iter, group, ex);
}
// To save memory
@@ -307,6 +312,7 @@ public class Learner {
// evaluation.add(LexiconFn.lexEval);
evaluation.logStats(prefix);
evaluation.putOutput(prefix);
+ evaluation.putOutput(prefix.replaceAll("iter=", "").replace('.', '_'));
}
private void printLearnerEventsIter(Example ex, int iter, String group) {
diff --git a/src/edu/stanford/nlp/sempre/NameValue.java b/src/edu/stanford/nlp/sempre/NameValue.java
index 3a82556..38d99f9 100644
--- a/src/edu/stanford/nlp/sempre/NameValue.java
+++ b/src/edu/stanford/nlp/sempre/NameValue.java
@@ -42,6 +42,7 @@ public class NameValue extends Value {
}
@Override public String sortString() { return id; }
+ @Override public String pureString() { return description == null ? id : description; }
@Override public int hashCode() { return id.hashCode(); }
@Override public boolean equals(Object o) {
diff --git a/src/edu/stanford/nlp/sempre/NumberValue.java b/src/edu/stanford/nlp/sempre/NumberValue.java
index 6408f1b..cd48725 100644
--- a/src/edu/stanford/nlp/sempre/NumberValue.java
+++ b/src/edu/stanford/nlp/sempre/NumberValue.java
@@ -40,6 +40,7 @@ public class NumberValue extends Value {
}
@Override public String sortString() { return "" + value; }
+ @Override public String pureString() { return "" + value; }
@Override public int hashCode() { return Double.valueOf(value).hashCode(); }
@Override public boolean equals(Object o) {
diff --git a/src/edu/stanford/nlp/sempre/Parser.java b/src/edu/stanford/nlp/sempre/Parser.java
index c3a7fa3..00fa76b 100644
--- a/src/edu/stanford/nlp/sempre/Parser.java
+++ b/src/edu/stanford/nlp/sempre/Parser.java
@@ -56,7 +56,7 @@ public abstract class Parser {
@Option(gloss = "Dump all features (for debugging)")
public boolean dumpAllFeatures = false;
-
+
@Option(gloss = "Call SetEvaluation during parsing")
public boolean callSetEvaluation = true;
}
@@ -100,8 +100,8 @@ public abstract class Parser {
this.valueEvaluator = spec.valueEvaluator;
computeCatUnaryRules();
- LogInfo.logs("Parser: %d catUnaryRules (sorted), %d nonCatUnaryRules (in trie)",
- catUnaryRules.size(), grammar.rules.size() - catUnaryRules.size());
+ LogInfo.logs("%s: %d catUnaryRules (sorted), %d nonCatUnaryRules (in trie)",
+ this.getClass().getSimpleName(), catUnaryRules.size(), grammar.rules.size() - catUnaryRules.size());
}
// If grammar changes, then we might need to update aspects of the parser.
@@ -210,6 +210,7 @@ public abstract class Parser {
int correctIndexAfterParse = -1;
double maxCompatibility = 0.0;
double[] compatibilities = null;
+ int numCorrect = 0, numPartialCorrect = 0, numIncorrect = 0;
if (ex.targetValue != null) {
compatibilities = new double[numCandidates];
@@ -221,6 +222,14 @@ public abstract class Parser {
correctIndex = i;
// record maximum compatibility for partial oracle
maxCompatibility = Math.max(compatibilities[i], maxCompatibility);
+ // Count
+ if (compatibilities[i] == 1) {
+ numCorrect++;
+ } else if (compatibilities[i] == 0) {
+ numIncorrect++;
+ } else {
+ numPartialCorrect++;
+ }
}
// What if we only had parsed bottom up?
for (int i = 0; i < numCandidates; i++) {
@@ -334,6 +343,9 @@ public abstract class Parser {
evaluation.add("numCandidates", numCandidates); // From this parse
if (numCandidates > 0)
evaluation.add("parsedNumCandidates", numCandidates);
+ evaluation.add("numCorrect", numCorrect);
+ evaluation.add("numPartialCorrect", numPartialCorrect);
+ evaluation.add("numIncorrect", numIncorrect);
// Add parsing stats
evaluation.add(state.evaluation);
diff --git a/src/edu/stanford/nlp/sempre/ParserState.java b/src/edu/stanford/nlp/sempre/ParserState.java
index 680fd40..9194bcb 100644
--- a/src/edu/stanford/nlp/sempre/ParserState.java
+++ b/src/edu/stanford/nlp/sempre/ParserState.java
@@ -24,7 +24,7 @@ public abstract class ParserState {
}
public static Options opts = new Options();
- public enum CustomExpectedCount { NONE, UNIFORM, TOP, RANDOM, }
+ public enum CustomExpectedCount { NONE, UNIFORM, TOP, TOPALT, RANDOM, }
//// Input: specification of how to parse
@@ -130,7 +130,7 @@ public abstract class ParserState {
for (Derivation deriv : derivations)
deriv.score += Parser.opts.derivationScoreRandom.nextDouble() * Parser.opts.derivationScoreNoise;
}
-
+
Derivation.sortByScore(derivations);
// Print out information
@@ -275,7 +275,7 @@ public abstract class ParserState {
predScores = new double[n];
// For update schemas that choose one good and one bad candidate to update
int[] goodAndBad = null;
- if (opts.customExpectedCounts == CustomExpectedCount.TOP) {
+ if (opts.customExpectedCounts == CustomExpectedCount.TOP || opts.customExpectedCounts == CustomExpectedCount.TOPALT) {
goodAndBad = getTopDerivations(derivations);
if (goodAndBad == null) return;
} else if (opts.customExpectedCounts == CustomExpectedCount.RANDOM) {
@@ -300,6 +300,10 @@ public abstract class ParserState {
trueScores[i] = (i == goodAndBad[0]) ? 0 : Double.NEGATIVE_INFINITY;
predScores[i] = (i == goodAndBad[1]) ? 0 : Double.NEGATIVE_INFINITY;
break;
+ case TOPALT:
+ trueScores[i] = (i == goodAndBad[0]) ? 0 : Double.NEGATIVE_INFINITY;
+ predScores[i] = (i == goodAndBad[0] || i == goodAndBad[1]) ? deriv.score : Double.NEGATIVE_INFINITY;
+ break;
default:
throw new RuntimeException("Unknown customExpectedCounts: " + opts.customExpectedCounts);
}
diff --git a/src/edu/stanford/nlp/sempre/Rule.java b/src/edu/stanford/nlp/sempre/Rule.java
index a871e0e..31f30b7 100644
--- a/src/edu/stanford/nlp/sempre/Rule.java
+++ b/src/edu/stanford/nlp/sempre/Rule.java
@@ -31,7 +31,7 @@ public class Rule {
public final SemanticFn sem; // Takes derivations corresponding to RHS categories and produces a set of derivations corresponding to LHS.
public List> info; // Extra info
public RuleSource source = null; // for tracking where the rule comes from when they are induced
-
+
// Cache the semanticRepn
public String getSemRepn() {
if (semRepn == null) semRepn = sem.getClass().getSimpleName();
@@ -51,9 +51,9 @@ public class Rule {
public String toString() {
if (stringRepn == null) {
String semStr = sem == null? "NullSemanticFn" : sem.toString();
- int maxLength = 100;
- if (semStr.length() > maxLength)
- semStr = String.format("%s...(%d total)", semStr.substring(0,maxLength), semStr.length());
+ //int maxLength = 100;
+ //if (semStr.length() > maxLength)
+ // semStr = String.format("%s...(%d total)", semStr.substring(0,maxLength), semStr.length());
stringRepn = lhs + " -> " + (rhs == null ? "" : Joiner.on(' ').join(rhs)) + " " + semStr;
}
return stringRepn;
@@ -139,13 +139,13 @@ public class Rule {
else
return f == 1.0 ? false : !FloatingParser.opts.defaultIsFloating;
}
-
+
public boolean isInduced() {
double a = getInfoTag("induced");
if (a == 1.0) return true;
return false;
}
-
+
@Override
public boolean equals(Object o) {
if (!(o instanceof Rule)) return false;
@@ -155,7 +155,7 @@ public class Rule {
public int hashCode() {
return this.toString().hashCode();
}
-
+
public String toJson() {
Map jsonMap = new LinkedHashMap<>();
jsonMap.put("lhs", lhs);
diff --git a/src/edu/stanford/nlp/sempre/StringValue.java b/src/edu/stanford/nlp/sempre/StringValue.java
index b67c426..91ebda6 100644
--- a/src/edu/stanford/nlp/sempre/StringValue.java
+++ b/src/edu/stanford/nlp/sempre/StringValue.java
@@ -20,6 +20,7 @@ public class StringValue extends Value {
}
@Override public String sortString() { return "\"" + value + "\""; }
+ @Override public String pureString() { return value; }
@Override public int hashCode() { return value.hashCode(); }
@Override public boolean equals(Object o) {
diff --git a/src/edu/stanford/nlp/sempre/TargetValuePreprocessor.java b/src/edu/stanford/nlp/sempre/TargetValuePreprocessor.java
index ac1c817..40d3938 100644
--- a/src/edu/stanford/nlp/sempre/TargetValuePreprocessor.java
+++ b/src/edu/stanford/nlp/sempre/TargetValuePreprocessor.java
@@ -27,10 +27,10 @@ public abstract class TargetValuePreprocessor {
}
public static void setSingleton(TargetValuePreprocessor processor) { singleton = processor; }
- public abstract Value preprocess(Value value);
+ public abstract Value preprocess(Value value, Example ex);
}
class IdentityTargetValuePreprocessor extends TargetValuePreprocessor {
- public Value preprocess(Value value) { return value; }
+ public Value preprocess(Value value, Example ex) { return value; }
}
diff --git a/src/edu/stanford/nlp/sempre/UriValue.java b/src/edu/stanford/nlp/sempre/UriValue.java
index baec2d6..e039898 100644
--- a/src/edu/stanford/nlp/sempre/UriValue.java
+++ b/src/edu/stanford/nlp/sempre/UriValue.java
@@ -20,6 +20,9 @@ public class UriValue extends Value {
return tree;
}
+ @Override public String sortString() { return "" + value; }
+ @Override public String pureString() { return "" + value; }
+
@Override public int hashCode() { return value.hashCode(); }
@Override public boolean equals(Object o) {
if (this == o) return true;
diff --git a/src/edu/stanford/nlp/sempre/Value.java b/src/edu/stanford/nlp/sempre/Value.java
index a69aa2a..486d5ec 100644
--- a/src/edu/stanford/nlp/sempre/Value.java
+++ b/src/edu/stanford/nlp/sempre/Value.java
@@ -24,6 +24,9 @@ public abstract class Value {
// (optional) String used for sorting Values. The default is to call toString()
public String sortString() { return toString(); }
+ // (optional) String without the LispTree structure. The default is to call toString()
+ public String pureString() { return toString(); }
+
@JsonCreator
public static Value fromString(String str) {
return Values.fromLispTree(LispTree.proto.parseFromString(str));
diff --git a/src/edu/stanford/nlp/sempre/cprune/CPruneDerivInfo.java b/src/edu/stanford/nlp/sempre/cprune/CPruneDerivInfo.java
new file mode 100644
index 0000000..3c487bb
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/CPruneDerivInfo.java
@@ -0,0 +1,13 @@
+package edu.stanford.nlp.sempre.cprune;
+
+import java.util.Map;
+import java.util.List;
+
+public class CPruneDerivInfo {
+
+ public Map treeSymbols;
+ public Map ruleSymbols;
+ public List customRuleStrings;
+ public boolean containsCrossReference;
+
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/CPruneFloatingParser.java b/src/edu/stanford/nlp/sempre/cprune/CPruneFloatingParser.java
new file mode 100644
index 0000000..2b6ef29
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/CPruneFloatingParser.java
@@ -0,0 +1,119 @@
+package edu.stanford.nlp.sempre.cprune;
+
+import java.util.List;
+
+import edu.stanford.nlp.sempre.*;
+import fig.basic.LogInfo;
+
+/**
+ * A parser that first tries to exploit the macro grammar and only fall back to full search when needed.
+ */
+public class CPruneFloatingParser extends FloatingParser {
+
+ FloatingParser exploreParser;
+
+ public CPruneFloatingParser(Spec spec) {
+ super(spec);
+ exploreParser = new FloatingParser(spec).setEarlyStopping(true, CollaborativePruner.opts.maxDerivations);
+ }
+
+ @Override
+ public void onBeginDataGroup(int iter, int numIters, String group) {
+ if (CollaborativePruner.uidToCachedNeighbors == null) {
+ CollaborativePruner.customGrammar.init(grammar);
+ CollaborativePruner.loadNeighbors();
+ }
+ CollaborativePruner.stats.reset(iter + "." + group);
+ }
+
+ @Override
+ public ParserState newParserState(Params params, Example ex, boolean computeExpectedCounts) {
+ return new CPruneFloatingParserState(this, params, ex, computeExpectedCounts);
+ }
+
+}
+
+class CPruneFloatingParserState extends ParserState {
+
+ public CPruneFloatingParserState(Parser parser, Params params, Example ex, boolean computeExpectedCounts) {
+ super(parser, params, ex, computeExpectedCounts);
+ }
+
+ @Override
+ public void infer() {
+ LogInfo.begin_track("CPruneFloatingParser.infer()");
+ boolean exploitSucceeds = exploit();
+ if (computeExpectedCounts) {
+ LogInfo.begin_track("Summary of Collaborative Pruning");
+ LogInfo.logs("Exploit succeeds: " + exploitSucceeds);
+ LogInfo.logs("Exploit success rate: " + CollaborativePruner.stats.successfulExploit + "/" + CollaborativePruner.stats.totalExploit);
+ LogInfo.end_track();
+ }
+ // Explore only on the first training iteration
+ if (CollaborativePruner.stats.iter.equals("0.train") && computeExpectedCounts && !exploitSucceeds
+ && (CollaborativePruner.stats.totalExplore <= CollaborativePruner.opts.maxExplorationIters)) {
+ explore();
+ LogInfo.logs("Consistent pattern: " + CollaborativePruner.getConsistentPattern(ex));
+ LogInfo.logs("Explore success rate: " + CollaborativePruner.stats.successfulExplore + "/" + CollaborativePruner.stats.totalExplore);
+ }
+ LogInfo.end_track();
+ }
+
+ public void explore() {
+ LogInfo.begin_track("Explore");
+ CollaborativePruner.initialize(ex, CollaborativePruner.Mode.EXPLORE);
+ ParserState exploreParserState = ((CPruneFloatingParser) parser).exploreParser.newParserState(params, ex, computeExpectedCounts);
+ exploreParserState.infer();
+ predDerivations.clear();
+ predDerivations.addAll(exploreParserState.predDerivations);
+ expectedCounts = exploreParserState.expectedCounts;
+ if (computeExpectedCounts) {
+ for (Derivation deriv : predDerivations)
+ CollaborativePruner.updateConsistentPattern(parser.valueEvaluator, ex, deriv);
+ }
+ CollaborativePruner.stats.totalExplore += 1;
+ if (CollaborativePruner.foundConsistentDerivation)
+ CollaborativePruner.stats.successfulExplore += 1;
+ LogInfo.end_track();
+ }
+
+ public boolean exploit() {
+ LogInfo.begin_track("Exploit");
+ CollaborativePruner.initialize(ex, CollaborativePruner.Mode.EXPLOIT);
+ Grammar miniGrammar = new MiniGrammar(CollaborativePruner.predictedRules);
+ Parser exploitParser = new FloatingParser(new Parser.Spec(miniGrammar, parser.extractor, parser.executor, parser.valueEvaluator));
+ ParserState exploitParserState = exploitParser.newParserState(params, ex, computeExpectedCounts);
+ exploitParserState.infer();
+ predDerivations.clear();
+ predDerivations.addAll(exploitParserState.predDerivations);
+ expectedCounts = exploitParserState.expectedCounts;
+ if (computeExpectedCounts) {
+ for (Derivation deriv : predDerivations)
+ CollaborativePruner.updateConsistentPattern(parser.valueEvaluator, ex, deriv);
+ }
+ boolean succeeds = CollaborativePruner.foundConsistentDerivation;
+ CollaborativePruner.stats.totalExploit += 1;
+ if (succeeds)
+ CollaborativePruner.stats.successfulExploit += 1;
+ LogInfo.end_track();
+ return succeeds;
+ }
+}
+
+// ============================================================
+// Helper classes
+// ============================================================
+
+class MiniGrammar extends Grammar {
+
+ public MiniGrammar(List rules) {
+ this.rules.addAll(rules);
+ if (CollaborativePruner.opts.verbose >= 2) {
+ LogInfo.begin_track("MiniGrammar Rules");
+ for (Rule rule : rules)
+ LogInfo.logs("%s %s", rule, rule.isAnchored() ? "[A]" : "[F]");
+ LogInfo.end_track();
+ }
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/CPruneStats.java b/src/edu/stanford/nlp/sempre/cprune/CPruneStats.java
new file mode 100644
index 0000000..6892bc3
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/CPruneStats.java
@@ -0,0 +1,20 @@
+package edu.stanford.nlp.sempre.cprune;
+
+/**
+ * Stores various statistic.
+ */
+public class CPruneStats {
+ public String iter;
+ public int totalExplore = 0;
+ public int successfulExplore = 0;
+ public int totalExploit = 0;
+ public int successfulExploit = 0;
+
+ public void reset(String iter) {
+ this.iter = iter;
+ this.totalExplore = 0;
+ this.successfulExplore = 0;
+ this.totalExploit = 0;
+ this.successfulExploit = 0;
+ }
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/CollaborativePruner.java b/src/edu/stanford/nlp/sempre/cprune/CollaborativePruner.java
new file mode 100644
index 0000000..c65a1c9
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/CollaborativePruner.java
@@ -0,0 +1,187 @@
+package edu.stanford.nlp.sempre.cprune;
+
+import java.io.*;
+import java.util.*;
+
+import fig.basic.*;
+import edu.stanford.nlp.sempre.*;
+
+/**
+ * Static class for collaborative pruning.
+ */
+public class CollaborativePruner {
+ public static class Options {
+ @Option(gloss = "Logging verbosity")
+ public int verbose = 0;
+ @Option(gloss = "K = Maximum number of nearest-neighbor examples to consider (-1 to use all examples so far)")
+ public int maxNumNeighbors = -1;
+ @Option(gloss = "Load cached neighbors from this file")
+ public String neighborFilePath = null;
+ @Option(gloss = "Maximum number of matching patterns (default = use all patterns)")
+ public int maxPredictedPatterns = Integer.MAX_VALUE;
+ @Option(gloss = "Maximum number of derivations per example")
+ public int maxDerivations = 5000;
+ @Option(gloss = "Maximum number of times to fall back to exploration")
+ public int maxExplorationIters = Integer.MAX_VALUE;
+ }
+
+ public static Options opts = new Options();
+
+ public enum Mode { EXPLORE, EXPLOIT, NONE }
+
+ public static Mode mode = Mode.NONE;
+ public static CPruneStats stats = new CPruneStats();
+ public static CustomGrammar customGrammar = new CustomGrammar();
+
+ // Static class; do not instantiate
+ private CollaborativePruner() { throw new RuntimeException("Cannot instantiate CollaborativePruner"); }
+
+ // Global variables
+ // Nearest neighbors
+ static Map> uidToCachedNeighbors;
+ // uid => pattern
+ static Map consistentPattern = new HashMap<>();
+ // patternString => customRuleString
+ static Map> customRules = new HashMap<>();
+ // set of patternStrings
+ static Set allConsistentPatterns = new HashSet<>();
+
+ // Example-level variables
+ public static boolean foundConsistentDerivation = false;
+ public static Map predictedPatterns;
+ public static List predictedRules;
+
+ /**
+ * Read the cached neighbors file.
+ * Line Format: ex_id [tab] neighbor_id1,neighbor_id2,...
+ */
+ public static void loadNeighbors() {
+ if (opts.neighborFilePath == null) {
+ LogInfo.logs("neighborFilePath is null.");
+ return;
+ }
+ LogInfo.begin_track("Loading cached neighbors from %s", opts.neighborFilePath);
+ uidToCachedNeighbors = new HashMap<>();
+ try {
+ BufferedReader reader = IOUtils.openIn(opts.neighborFilePath);
+ String line;
+ while ((line = reader.readLine()) != null) {
+ String[] tokens = line.split("\t");
+ String uid = tokens[0];
+ String[] nids = tokens[1].split(",");
+ uidToCachedNeighbors.put(uid, Arrays.asList(nids));
+ }
+ reader.close();
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ LogInfo.end_track();
+ }
+
+ public static void initialize(Example ex, Mode mode) {
+ CollaborativePruner.mode = mode;
+ predictedRules = null;
+ predictedPatterns = null;
+ foundConsistentDerivation = false;
+ if (mode == Mode.EXPLOIT) {
+ preprocessExample(ex);
+ }
+ }
+
+ static void preprocessExample(Example ex) {
+ Map patternFreqMap = new HashMap<>();
+ List cachedNeighbors = uidToCachedNeighbors.get(ex.id);
+ int total = 0;
+
+ // Gather the neighbors
+ if (opts.maxNumNeighbors > 0) {
+ for (String nid : cachedNeighbors) {
+ // Only get examples that have been previously processed + found a consistent formula
+ if (!consistentPattern.containsKey(nid))
+ continue;
+
+ String neighborPattern = consistentPattern.get(nid).pattern;
+ if (!patternFreqMap.containsKey(neighborPattern))
+ patternFreqMap.put(neighborPattern, new FormulaPattern(neighborPattern, 0));
+ patternFreqMap.get(neighborPattern).frequency++;
+ total++;
+ if (total >= opts.maxNumNeighbors)
+ break;
+ }
+ } else {
+ for (String patternString : allConsistentPatterns) {
+ patternFreqMap.put(patternString, new FormulaPattern(patternString, 1));
+ }
+ }
+
+ // Sort by frequency (more frequent = smaller; see FormulaPattern.compareTo)
+ List> patternFreqEntries = new ArrayList<>(patternFreqMap.entrySet());
+ patternFreqEntries.sort(new ValueComparator<>(false));
+
+ // Gather the patterns
+ LogInfo.begin_track("Predicted patterns");
+ int rank = 0;
+ Set predictedRulesStrings = new HashSet<>();
+ predictedPatterns = new HashMap<>();
+ for (Map.Entry entry : patternFreqEntries) {
+ FormulaPattern newPattern = entry.getValue();
+ predictedPatterns.put(newPattern.pattern, newPattern);
+ predictedRulesStrings.addAll(customRules.get(newPattern.pattern));
+ LogInfo.logs((rank + 1) + ". " + newPattern.pattern + " (" + newPattern.frequency + ")");
+ rank++;
+ if (rank >= opts.maxPredictedPatterns)
+ break;
+ }
+ // Gather the rules
+ predictedRules = customGrammar.getRules(predictedRulesStrings);
+ LogInfo.end_track();
+ }
+
+ public static String getPatternString(Derivation deriv) {
+ if (deriv.cat.equals("$TOKEN") || deriv.cat.equals("$PHRASE")
+ || deriv.cat.equals("$LEMMA_TOKEN") || deriv.cat.equals("$LEMMA_PHRASE")) {
+ return deriv.cat;
+ } else {
+ return FormulaPattern.convertToIndexedPattern(deriv);
+ }
+ }
+
+ public static void addRules(String patternString, Derivation deriv, Example ex) {
+ if (!customRules.containsKey(patternString)) {
+ customRules.put(patternString, new HashSet());
+ }
+ Set parsedCustomRules = customGrammar.addCustomRule(deriv, ex);
+ customRules.get(patternString).addAll(parsedCustomRules);
+ }
+
+ /**
+ * Get called when a (consistent) formula is found.
+ * Update the consistent patterns.
+ */
+ public static void updateConsistentPattern(ValueEvaluator evaluator, Example ex, Derivation deriv) {
+ String uid = ex.id;
+ if (ex.targetValue != null)
+ deriv.compatibility = evaluator.getCompatibility(ex.targetValue, deriv.value);
+
+ if (deriv.isRootCat() && deriv.compatibility == 1) {
+ foundConsistentDerivation = true;
+ LogInfo.logs("Found consistent deriv: %s", deriv);
+
+ String patternString = getPatternString(deriv);
+ FormulaPattern newConsistentPattern = new FormulaPattern(patternString, 0);
+ newConsistentPattern.score = deriv.getScore();
+
+ FormulaPattern oldConsistentPattern = consistentPattern.get(uid);
+ if (oldConsistentPattern == null || newConsistentPattern.score > oldConsistentPattern.score) {
+ addRules(patternString, deriv, ex);
+ consistentPattern.put(uid, newConsistentPattern);
+ allConsistentPatterns.add(patternString);
+ }
+ }
+ }
+
+ public static FormulaPattern getConsistentPattern(Example ex) {
+ return consistentPattern.get(ex.id);
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/CustomGrammar.java b/src/edu/stanford/nlp/sempre/cprune/CustomGrammar.java
new file mode 100644
index 0000000..1dbca83
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/CustomGrammar.java
@@ -0,0 +1,268 @@
+package edu.stanford.nlp.sempre.cprune;
+
+import java.util.*;
+
+import edu.stanford.nlp.sempre.*;
+import fig.basic.*;
+
+public class CustomGrammar extends Grammar {
+ public static class Options {
+ @Option(gloss = "Whether to decompose the templates into multiple rules")
+ public boolean enableTemplateDecomposition = true;
+ }
+
+ public static Options opts = new Options();
+
+ public static final Set baseCategories = new HashSet(Arrays.asList(
+ Rule.tokenCat, Rule.phraseCat, Rule.lemmaTokenCat, Rule.lemmaPhraseCat,
+ "$Unary", "$Binary", "$Entity", "$Property"));
+
+ ArrayList baseRules = new ArrayList<>();
+ // symbolicFormulas => symbolicFormula ID
+ Map symbolicFormulas = new HashMap<>();
+ // indexedSymbolicFormula => customRuleString
+ Map> customRules = new HashMap<>();
+ // customRuleString => Binarized rules
+ Map> customBinarizedRules = new HashMap<>();
+
+ public void init(Grammar initGrammar) {
+ baseRules = new ArrayList<>();
+ for (Rule rule : initGrammar.getRules()) {
+ if (baseCategories.contains(rule.lhs)) {
+ baseRules.add(rule);
+ }
+ }
+ this.freshCatIndex = initGrammar.getFreshCatIndex();
+ }
+
+ public List getRules(Collection customRuleStrings) {
+ Set ruleSet = new LinkedHashSet<>();
+ ruleSet.addAll(baseRules);
+ for (String ruleString : customRuleStrings) {
+ ruleSet.addAll(customBinarizedRules.get(ruleString));
+ }
+ return new ArrayList(ruleSet);
+ }
+
+ public Set addCustomRule(Derivation deriv, Example ex) {
+ String indexedSymbolicFormula = getIndexedSymbolicFormula(deriv);
+ if (customRules.containsKey(indexedSymbolicFormula)) {
+ return customRules.get(indexedSymbolicFormula);
+ }
+
+ CPruneDerivInfo derivInfo = aggregateSymbols(deriv);
+ Set crossReferences = new HashSet<>();
+ for (Symbol symbol : derivInfo.treeSymbols.values()) {
+ if (symbol.frequency > 1) {
+ crossReferences.add(symbol.formula);
+ }
+ }
+ computeCustomRules(deriv, crossReferences);
+ customRules.put(indexedSymbolicFormula, new HashSet(derivInfo.customRuleStrings));
+
+ LogInfo.begin_track("Add custom rules for formula: " + indexedSymbolicFormula);
+ for (String customRuleString : derivInfo.customRuleStrings) {
+ if (customBinarizedRules.containsKey(customRuleString)) {
+ LogInfo.log("Custom rule exists: " + customRuleString);
+ continue;
+ }
+
+ rules = new ArrayList<>();
+ LispTree tree = LispTree.proto.parseFromString(customRuleString);
+ interpretRule(tree);
+ customBinarizedRules.put(customRuleString, new HashSet(rules));
+
+ // Debug
+ LogInfo.begin_track("Add custom rule: " + customRuleString);
+ for (Rule rule : rules) {
+ LogInfo.log(rule.toString());
+ }
+ LogInfo.end_track();
+ }
+ LogInfo.end_track();
+
+ // Debug
+ System.out.println("consistent_lf\t" + ex.id + "\t" + deriv.formula.toString());
+
+ return customRules.get(indexedSymbolicFormula);
+ }
+
+ public static String getIndexedSymbolicFormula(Derivation deriv) {
+ return getIndexedSymbolicFormula(deriv, deriv.formula.toString());
+ }
+
+ /**
+ * Replace symbols (e.g., fb:row.row.name) with placeholders (e.g., Binary#1).
+ */
+ public static String getIndexedSymbolicFormula(Derivation deriv, String formula) {
+ CPruneDerivInfo derivInfo = aggregateSymbols(deriv);
+ int index = 1;
+ List symbolList = new ArrayList<>(derivInfo.treeSymbols.values());
+ for (Symbol symbol : symbolList)
+ symbol.computeIndex(formula);
+ Collections.sort(symbolList);
+ for (Symbol symbol : symbolList) {
+ if (formula.equals(symbol.formula))
+ formula = symbol.category + "#" + index;
+ formula = safeReplace(formula, symbol.formula, symbol.category + "#" + index);
+ index += 1;
+ }
+ return formula;
+ }
+
+ // ============================================================
+ // Private methods
+ // ============================================================
+
+ private static String safeReplace(String formula, String target, String replacement) {
+ // (argmin 1 1 ...) and (argmax 1 1 ...) are troublesome
+ String before = formula, targetBefore = target;
+ formula = formula.replace("(argmin (number 1) (number 1)", "(ARGMIN");
+ formula = formula.replace("(argmax (number 1) (number 1)", "(ARGMAX");
+ target = target.replace("(argmin (number 1) (number 1)", "(ARGMIN");
+ target = target.replace("(argmax (number 1) (number 1)", "(ARGMAX");
+ formula = formula.replace(target + ")", replacement + ")");
+ formula = formula.replace(target + " ", replacement + " ");
+ formula = formula.replace("(ARGMIN", "(argmin (number 1) (number 1)");
+ formula = formula.replace("(ARGMAX", "(argmax (number 1) (number 1)");
+ if (CollaborativePruner.opts.verbose >= 2)
+ LogInfo.logs("REPLACE: [%s | %s] %s | %s", targetBefore, replacement, before, formula);
+ return formula;
+ }
+
+ /**
+ * Cache the symbols in deriv.tempState[cprune].treeSymbols
+ */
+ private static CPruneDerivInfo aggregateSymbols(Derivation deriv) {
+ Map tempState = deriv.getTempState();
+ if (tempState.containsKey("cprune")) {
+ return (CPruneDerivInfo) tempState.get("cprune");
+ }
+ CPruneDerivInfo derivInfo = new CPruneDerivInfo();
+ tempState.put("cprune", derivInfo);
+
+ Map treeSymbols = new LinkedHashMap<>();
+ derivInfo.treeSymbols = treeSymbols;
+ if (baseCategories.contains(deriv.cat)) {
+ String formula = deriv.formula.toString();
+ treeSymbols.put(formula, new Symbol(deriv.cat, formula, 1));
+ } else {
+ for (Derivation child : deriv.children) {
+ CPruneDerivInfo childInfo = aggregateSymbols(child);
+ for (Symbol symbol : childInfo.treeSymbols.values()) {
+ if (derivInfo.treeSymbols.containsKey(symbol.formula)) {
+ treeSymbols.get(symbol.formula).frequency += symbol.frequency;
+ } else {
+ treeSymbols.put(symbol.formula, symbol);
+ }
+ }
+ }
+ }
+ return derivInfo;
+ }
+
+ private CPruneDerivInfo computeCustomRules(Derivation deriv, Set crossReferences) {
+ CPruneDerivInfo derivInfo = (CPruneDerivInfo) deriv.getTempState().get("cprune");
+ Map ruleSymbols = new LinkedHashMap<>();
+ derivInfo.ruleSymbols = ruleSymbols;
+ derivInfo.customRuleStrings = new ArrayList<>();
+ String formula = deriv.formula.toString();
+
+ if (baseCategories.contains(deriv.cat)) {
+ // Leaf node induces no custom rule
+ derivInfo.containsCrossReference = crossReferences.contains(formula);
+ // Propagate the symbol of this derivation to the parent
+ ruleSymbols.putAll(derivInfo.treeSymbols);
+ } else {
+ derivInfo.containsCrossReference = false;
+ for (Derivation child : deriv.children) {
+ CPruneDerivInfo childInfo = computeCustomRules(child, crossReferences);
+ derivInfo.containsCrossReference = derivInfo.containsCrossReference || childInfo.containsCrossReference;
+ }
+
+ for (Derivation child : deriv.children) {
+ CPruneDerivInfo childInfo = (CPruneDerivInfo) child.getTempState().get("cprune");
+ ruleSymbols.putAll(childInfo.ruleSymbols);
+ derivInfo.customRuleStrings.addAll(childInfo.customRuleStrings);
+ }
+
+ if (opts.enableTemplateDecomposition == false || derivInfo.containsCrossReference) {
+ // If this node contains a cross reference
+ if (deriv.isRootCat()) {
+ // If this is the root node, then generate a custom rule
+ derivInfo.customRuleStrings.add(getCustomRuleString(deriv, derivInfo));
+ }
+ } else {
+ if (!deriv.cat.startsWith("$Intermediate")) {
+ // Generate a custom rule for this node
+ derivInfo.customRuleStrings.add(getCustomRuleString(deriv, derivInfo));
+
+ // Propagate this derivation as a category to the parent
+ ruleSymbols.clear();
+ ruleSymbols.put(formula, new Symbol(hash(deriv), deriv.formula.toString(), 1));
+ }
+ }
+ }
+ return derivInfo;
+ }
+
+ private String getCustomRuleString(Derivation deriv, CPruneDerivInfo derivInfo) {
+ String formula = deriv.formula.toString();
+ List rhsSymbols = new ArrayList<>(derivInfo.ruleSymbols.values());
+ for (Symbol symbol : rhsSymbols)
+ symbol.computeIndex(formula);
+ Collections.sort(rhsSymbols);
+
+ String lhs = null;
+ if (derivInfo.containsCrossReference)
+ lhs = deriv.cat;
+ else
+ lhs = deriv.isRootCat() ? "$ROOT" : hash(deriv);
+
+ LinkedList rhsList = new LinkedList<>();
+ int index = 1;
+ for (Symbol symbol : rhsSymbols) {
+ if (formula.equals(symbol.formula)) {
+ formula = "(IdentityFn)";
+ } else {
+ formula = safeReplace(formula, symbol.formula, "(var s" + index + ")");
+ formula = "(lambda s" + index + " " + formula + ")";
+ }
+ rhsList.addFirst(symbol.category);
+ index += 1;
+ }
+ String rhs = null;
+ if (rhsList.size() > 0) {
+ rhs = "(" + String.join(" ", rhsList) + ")";
+ } else {
+ rhs = "(nothing)";
+ formula = "(ConstantFn " + formula + ")";
+ }
+ return "(rule " + lhs + " " + rhs + " " + formula + ")";
+ }
+
+ private String hash(Derivation deriv) {
+ if (baseCategories.contains(deriv.cat))
+ return deriv.cat;
+
+ String formula = getSymbolicFormula(deriv);
+ if (!symbolicFormulas.containsKey(formula)) {
+ symbolicFormulas.put(formula, symbolicFormulas.size() + 1);
+ String hashString = "$Formula" + symbolicFormulas.get(formula);
+ LogInfo.log("Add symbolic formula: " + hashString + " = " + formula + " (" + deriv.cat + ")");
+ }
+ return "$Formula" + symbolicFormulas.get(formula);
+ }
+
+ private static String getSymbolicFormula(Derivation deriv) {
+ CPruneDerivInfo derivInfo = aggregateSymbols(deriv);
+ String formula = deriv.formula.toString();
+ for (Symbol symbol : derivInfo.treeSymbols.values()) {
+ if (formula.equals(symbol.formula))
+ formula = symbol.category;
+ formula = safeReplace(formula, symbol.formula, symbol.category);
+ }
+ return formula;
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/FormulaPattern.java b/src/edu/stanford/nlp/sempre/cprune/FormulaPattern.java
new file mode 100644
index 0000000..0c83d04
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/FormulaPattern.java
@@ -0,0 +1,106 @@
+package edu.stanford.nlp.sempre.cprune;
+
+import java.util.regex.Pattern;
+
+import edu.stanford.nlp.sempre.Derivation;
+import fig.basic.LogInfo;
+
+public class FormulaPattern implements Comparable {
+ public String pattern;
+ public Integer frequency;
+ public Double score;
+
+ public FormulaPattern(String pattern, Integer frequency) {
+ this.pattern = pattern;
+ this.frequency = frequency;
+ }
+
+ public Double complexity() {
+ // Roughly the number of predicates
+ return (double) (pattern.length() - pattern.replace("(@R", "***").replace("(", "").length());
+ }
+
+ @Override
+ public String toString() {
+ return "(" + pattern + ", " + frequency + ")";
+ }
+
+ @Override
+ public int compareTo(FormulaPattern that) {
+ if (this.frequency > that.frequency) {
+ return -1;
+ } else if (this.frequency < that.frequency) {
+ return 1;
+ } else {
+ return this.complexity().compareTo(that.complexity());
+ }
+ }
+
+ // ============================================================
+ // Utilities
+ // ============================================================
+
+ private static Pattern reverseRelation = Pattern.compile("!(fb:[._a-z0-9]+)");
+ private static Pattern varName = Pattern.compile("\\((lambda|var) [a-z0-9]+");
+ private static Pattern compare = Pattern.compile("(<=|>=|>|<)");
+ private static Pattern whitespace = Pattern.compile("\\s+");
+
+ public static String convertToIndexedPattern(Derivation deriv) {
+ String formula = deriv.formula.toString();
+
+ // These can interfere with (number 1)
+ formula = formula.replace("argmax (number 1) (number 1)", "argmax");
+ formula = formula.replace("argmin (number 1) (number 1)", "argmin");
+
+ formula = removePropertyPredicates(formula);
+ formula = CustomGrammar.getIndexedSymbolicFormula(deriv, formula);
+
+ formula = formula.replace("fb:type.object.type fb:type.row", "@type @row");
+ formula = reverseRelation.matcher(formula).replaceAll("(reverse $1)");
+ formula = formula.replace("fb:row.row.index", "(reverse (lambda x ((reverse @index) (var x))))");
+ formula = formula.replace("fb:row.row.next", "@next");
+ formula = varName.matcher(formula).replaceAll("($1 x");
+ formula = formula.replace("reverse", "@R");
+ formula = compare.matcher(formula).replaceAll("@compare");
+ formula = whitespace.matcher(formula).replaceAll(" ");
+
+ if (CollaborativePruner.opts.verbose >= 2)
+ LogInfo.logs("PATTERN: %s -> %s", deriv.formula, formula);
+ return formula;
+ }
+
+ private static Pattern cellProperty = Pattern.compile("!?fb:cell\\.cell\\.[_a-z0-9]+|\\(reverse fb:cell\\.cell\\.[_a-z0-9]+\\)");
+
+ /**
+ * Remove cell property relations (fb:cell.cell.*)
+ */
+ public static String removePropertyPredicates(String formula) {
+ formula = cellProperty.matcher(formula).replaceAll("@PPT");
+ while (formula.contains("@PPT")) {
+ int begin = formula.indexOf("(@PPT");
+ if (begin == -1) {
+ formula = formula.replace("@PPT", "");
+ break;
+ }
+ // Find the matching parenthesis
+ int count = 1;
+ for (int i = begin + 1; i < formula.length(); i++) {
+ if (formula.charAt(i) == '(') {
+ count++;
+ } else if (formula.charAt(i) == ')') {
+ count--;
+ if (count == 0) {
+ int end = i;
+ formula = formula.substring(0, begin) + formula.substring(begin + 6, end) + formula.substring(end + 1, formula.length());
+ break;
+ }
+ }
+ if (i == formula.length() - 1) {
+ LogInfo.fails("Unbalanced parentheses: %s", formula);
+ }
+ }
+ }
+ return formula;
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/cprune/Symbol.java b/src/edu/stanford/nlp/sempre/cprune/Symbol.java
new file mode 100644
index 0000000..75d8537
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/cprune/Symbol.java
@@ -0,0 +1,31 @@
+package edu.stanford.nlp.sempre.cprune;
+
+/**
+ * Represents the leaf node of the parse tree.
+ *
+ * Any sub-derivation whose category is in CustomGrammar.baseCategories becomes a Symbol.
+ */
+public class Symbol implements Comparable {
+ String category;
+ String formula;
+ Integer frequency;
+ Integer index;
+
+ public Symbol(String category, String formula, int frequency) {
+ this.category = category;
+ this.formula = formula;
+ this.frequency = frequency;
+ }
+
+ public void computeIndex(String referenceString) {
+ index = referenceString.indexOf(formula);
+ if (index < 0) {
+ index = Integer.MAX_VALUE;
+ }
+ }
+
+ @Override
+ public int compareTo(Symbol that) {
+ return index.compareTo(that.index);
+ }
+}
diff --git a/src/edu/stanford/nlp/sempre/tables/StringNormalizationUtils.java b/src/edu/stanford/nlp/sempre/tables/StringNormalizationUtils.java
index 3093007..fbb6eb7 100644
--- a/src/edu/stanford/nlp/sempre/tables/StringNormalizationUtils.java
+++ b/src/edu/stanford/nlp/sempre/tables/StringNormalizationUtils.java
@@ -53,6 +53,7 @@ public final class StringNormalizationUtils {
// Cell normalization
// ============================================================
+ public static final Pattern STRICT_DASH = Pattern.compile("\\s*[-‐‑⁃‒–—―]\\s*");
public static final Pattern DASH = Pattern.compile("\\s*[-‐‑⁃‒–—―/,:;]\\s*");
public static final Pattern COMMA = Pattern.compile("\\s*(,\\s|\\n|/)\\s*");
public static final Pattern SPACE = Pattern.compile("\\s+");
@@ -200,22 +201,31 @@ public final class StringNormalizationUtils {
return null;
}
- public static final DateTimeFormatter dateFormat = DateTimeFormat.forPattern("MMM d, yyyy");
+ public static final DateTimeFormatter americanDateFormat = DateTimeFormat.forPattern("MMM d, yyyy");
+ public static final Pattern suTimeDateFormat = Pattern.compile("([0-9X]{4})(?:-([0-9X]{2}))?(?:-([0-9X]{2}))?");
/**
* Convert string to DateValue.
*/
public static DateValue parseDate(String s) {
+ Matcher matcher = suTimeDateFormat.matcher(s.toUpperCase());
+ if (matcher.matches()) {
+ String yS = matcher.group(1), mS = matcher.group(2), dS = matcher.group(3);
+ int y = -1, m = -1, d = -1;
+ if (!(yS == null || yS.isEmpty() || yS.contains("X"))) y = Integer.parseInt(yS);
+ if (!(mS == null || mS.isEmpty() || mS.contains("X"))) m = Integer.parseInt(mS);
+ if (!(dS == null || dS.isEmpty() || dS.contains("X"))) d = Integer.parseInt(dS);
+ if (y == -1 && m == -1 && d == -1) return null;
+ return new DateValue(y, m, d);
+ }
try {
- DateTime date = dateFormat.parseDateTime(s);
+ DateTime date = americanDateFormat.parseDateTime(s);
return new DateValue(date.getYear(), date.getMonthOfYear(), date.getDayOfMonth());
} catch (IllegalArgumentException e) {
return null;
}
}
- public static final Pattern suTimeDateFormat = Pattern.compile("([0-9X]{4})(?:-([0-9X]{2}))?(?:-([0-9X]{2}))?");
-
public static DateValue parseDateWithLanguageAnalyzer(LanguageInfo languageInfo) {
if (languageInfo.numTokens() == 0) return null;
String nerSpan = languageInfo.getNormalizedNerSpan("DATE", 0, languageInfo.numTokens());
@@ -307,8 +317,7 @@ public final class StringNormalizationUtils {
.replaceAll("[‘’´`]", "'")
.replaceAll("[“”«»]", "\"")
.replaceAll("[•†‡]", "")
- .replaceAll("[‐‑–—]", "-")
- .replaceAll("[\\u2E00-\\uFFFF]", ""); // (Sorry Chinese people)
+ .replaceAll("[-‐‑–—]", "-");
return string.replaceAll("\\s+", " ").trim();
}
@@ -320,13 +329,8 @@ public final class StringNormalizationUtils {
// Citation
string = string.replaceAll("\\[(nb ?)?\\d+\\]", "");
string = string.replaceAll("\\*+$", "");
- // Year in parentheses
- string = string.replaceAll("\\(\\d* ?-? ?\\d*\\)", "");
// Outside Quote
string = string.replaceAll("^\"(.*)\"$", "$1");
- // Numbering
- if (!string.matches("^[0-9.]+$"))
- string = string.replaceAll("^\\d+\\.", "");
return string.replaceAll("\\s+", " ").trim();
}
@@ -336,13 +340,49 @@ public final class StringNormalizationUtils {
public static String aggressiveNormalize(String string) {
// Dashed / Parenthesized information
string = simpleNormalize(string);
- string = string.replaceAll("\\[[^\\]]*\\]", "");
- string = string.replaceAll("[\\u007F-\\uFFFF]", "");
- string = string.trim().replaceAll(" - .*$", "");
- string = string.trim().replaceAll("\\([^)]*\\)$", "");
+ String oldString;
+ do {
+ oldString = string;
+ // Remove citations
+ string = string.trim().replaceAll("((? metadata = ArrayListMultimap.create();
TableColumn column = new TableColumn("Test", "test", 0);
analyzeString(o, metadata, column, new HashMap<>());
- LogInfo.logs("%s %s", o, metadata);
+ String aggressive = aggressiveNormalize(o).toLowerCase();
+ String official = officialEvaluatorNormalize(o);
+ LogInfo.logs("%s %s | %s %s %s", o, metadata, official, aggressive, aggressive.equals(official));
}
public static void main(String[] args) {
@@ -363,6 +405,8 @@ public final class StringNormalizationUtils {
unitTest("twenty three");
unitTest("apple, banana, banana, BANANA");
unitTest("apple\nbanana\norange");
+ unitTest("0-1\n(4-5 p)");
+ unitTest("\"HELLO\"");
unitTest("21st");
unitTest("2001st");
unitTest("2,000,000 ft.");
diff --git a/src/edu/stanford/nlp/sempre/tables/TableKnowledgeGraph.java b/src/edu/stanford/nlp/sempre/tables/TableKnowledgeGraph.java
index 01b31ad..9300eb3 100644
--- a/src/edu/stanford/nlp/sempre/tables/TableKnowledgeGraph.java
+++ b/src/edu/stanford/nlp/sempre/tables/TableKnowledgeGraph.java
@@ -65,6 +65,8 @@ public class TableKnowledgeGraph extends KnowledgeGraph implements FuzzyMatchabl
Map relationIdToTableColumn;
// "fb:cell.palo_alto_ca" --> TableCellProperties object
Map cellIdToTableCellProperties;
+ // "fb:part.palo_alto" --> String
+ Map partIdToOriginalString;
FuzzyMatcher fuzzyMatcher;
public ExecutorCache executorCache;
@@ -143,9 +145,14 @@ public class TableKnowledgeGraph extends KnowledgeGraph implements FuzzyMatchabl
// Collect cell properties for public access
cellProperties = new HashSet<>(cellIdToTableCellProperties.values());
cellParts = new HashSet<>();
- for (TableCellProperties properties : cellProperties)
- for (Value part : properties.metadata.get(TableTypeSystem.CELL_PART_VALUE))
- cellParts.add((NameValue) part);
+ partIdToOriginalString = new HashMap<>();
+ for (TableCellProperties properties : cellProperties) {
+ for (Value part : properties.metadata.get(TableTypeSystem.CELL_PART_VALUE)) {
+ NameValue partNameValue = (NameValue) part;
+ cellParts.add(partNameValue);
+ partIdToOriginalString.put(partNameValue.id, partNameValue.description);
+ }
+ }
// Precompute normalized strings for fuzzy matching
fuzzyMatcher = FuzzyMatcher.getFuzzyMatcher(this);
executorCache = opts.individualExecutorCache ? new ExecutorCache() : null;
@@ -615,6 +622,8 @@ public class TableKnowledgeGraph extends KnowledgeGraph implements FuzzyMatchabl
if (nameValueId.startsWith("!")) nameValueId = nameValueId.substring(1);
if (cellIdToTableCellProperties.containsKey(nameValueId))
return cellIdToTableCellProperties.get(nameValueId).originalString;
+ if (partIdToOriginalString.containsKey(nameValueId))
+ return partIdToOriginalString.get(nameValueId);
if (relationIdToTableColumn.containsKey(nameValueId))
return relationIdToTableColumn.get(nameValueId).originalString;
if (nameValueId.startsWith(TableTypeSystem.CELL_SPECIFIC_TYPE_PREFIX)) {
@@ -660,6 +669,18 @@ public class TableKnowledgeGraph extends KnowledgeGraph implements FuzzyMatchabl
return answer;
}
+ /**
+ * Return the index of the column with the specified ID. Return -1 if not found.
+ */
+ public int getColumnIndex(String nameValueId) {
+ if (nameValueId.startsWith("!"))
+ nameValueId = nameValueId.substring(1);
+ for (int j = 0; j < columns.size(); j++) {
+ if (columns.get(j).relationNameValue.id.equals(nameValueId)) return j;
+ }
+ return -1;
+ }
+
// ============================================================
// Test
// ============================================================
diff --git a/src/edu/stanford/nlp/sempre/tables/TableValueEvaluator.java b/src/edu/stanford/nlp/sempre/tables/TableValueEvaluator.java
index 87e44a5..c825c7d 100644
--- a/src/edu/stanford/nlp/sempre/tables/TableValueEvaluator.java
+++ b/src/edu/stanford/nlp/sempre/tables/TableValueEvaluator.java
@@ -25,6 +25,8 @@ public class TableValueEvaluator implements ValueEvaluator {
public boolean ignoreNumberValueUnits = true;
@Option(gloss = "Strict date evaluation (year, month, and date all have to match)")
public boolean strictDateEvaluation = false;
+ @Option(gloss = "Check if the normalized text matches the official evaluator")
+ public boolean checkStringNormalization = false;
}
public static Options opts = new Options();
@@ -69,8 +71,16 @@ public class TableValueEvaluator implements ValueEvaluator {
String predText = (pred instanceof NameValue) ? ((NameValue) pred).description : ((DescriptionValue) pred).value;
if (predText == null) predText = "";
if (opts.allowNormalizedStringMatch) {
- targetText = StringNormalizationUtils.aggressiveNormalize(targetText);
- predText = StringNormalizationUtils.aggressiveNormalize(predText);
+ targetText = StringNormalizationUtils.aggressiveNormalize(targetText).toLowerCase();
+ predText = StringNormalizationUtils.aggressiveNormalize(predText).toLowerCase();
+ if (opts.checkStringNormalization) {
+ String targetTextOfficial = StringNormalizationUtils.officialEvaluatorNormalize(targetText);
+ String predTextOfficial = StringNormalizationUtils.officialEvaluatorNormalize(predText);
+ if (!targetTextOfficial.equals(targetText) && !(targetTextOfficial + ".").equals(targetText))
+ LogInfo.warnings("Different normalization: [%s][%s]", targetTextOfficial, targetText);
+ if (!predTextOfficial.equals(predText) && !(predTextOfficial + ".").equals(predText))
+ LogInfo.warnings("Different normalization: [%s][%s]", predTextOfficial, predText);
+ }
}
return targetText.equals(predText);
} else if (pred instanceof NumberValue) {
diff --git a/src/edu/stanford/nlp/sempre/tables/TableValuePreprocessor.java b/src/edu/stanford/nlp/sempre/tables/TableValuePreprocessor.java
index 49b7eed..4e62bfc 100644
--- a/src/edu/stanford/nlp/sempre/tables/TableValuePreprocessor.java
+++ b/src/edu/stanford/nlp/sempre/tables/TableValuePreprocessor.java
@@ -1,5 +1,9 @@
package edu.stanford.nlp.sempre.tables;
+import java.io.BufferedReader;
+import java.io.File;
+import java.io.FileReader;
+import java.io.IOException;
import java.util.*;
import edu.stanford.nlp.sempre.*;
@@ -8,15 +12,20 @@ import fig.basic.*;
public class TableValuePreprocessor extends TargetValuePreprocessor {
public static class Options {
@Option(gloss = "Verbosity") public int verbose = 0;
+ @Option(gloss = "Read preprocessed values from these .tagged files")
+ public List taggedFiles = new ArrayList<>();
}
public static Options opts = new Options();
@Override
- public Value preprocess(Value value) {
+ public Value preprocess(Value value, Example ex) {
+ if (!opts.taggedFiles.isEmpty() && ex != null) {
+ return getFromTaggedFile(ex.id);
+ }
if (value instanceof ListValue) {
List values = new ArrayList<>();
for (Value entry : ((ListValue) value).values) {
- values.add(preprocess(entry));
+ values.add(preprocessSingle(entry));
}
return new ListValue(values);
} else {
@@ -67,4 +76,76 @@ public class TableValuePreprocessor extends TargetValuePreprocessor {
return new DescriptionValue(origString);
}
+ // ============================================================
+ // Get preprocessed value from tagged file
+ // ============================================================
+
+ Map idToValue = null;
+
+ public Value getFromTaggedFile(String id) {
+ if (idToValue == null) readTaggedFiles();
+ return idToValue.get(id);
+ }
+
+ protected void readTaggedFiles() {
+ LogInfo.begin_track("Reading .tagged files");
+ idToValue = new HashMap<>();
+ for (String path : opts.taggedFiles) {
+ File file = new File(path);
+ if (file.isDirectory()) {
+ for (File subpath : file.listFiles())
+ readTaggedFile(subpath.toString());
+ } else {
+ readTaggedFile(path);
+ }
+ }
+ LogInfo.logs("Read %d entries", idToValue.size());
+ LogInfo.end_track();
+ }
+
+ protected void readTaggedFile(String path) {
+ LogInfo.begin_track("Reading %s", path);
+ try (BufferedReader reader = new BufferedReader(new FileReader(path))) {
+ // Read header
+ String[] header = reader.readLine().split("\t", -1);
+ int exIdIndex = 0, targetCanonIndex = 0;
+ while (!"id".equals(header[exIdIndex]))
+ exIdIndex++;
+ while (!"targetCanon".equals(header[targetCanonIndex]))
+ targetCanonIndex++;
+ // Read each line
+ String line;
+ while ((line = reader.readLine()) != null) {
+ String[] fields = line.split("\t", -1); // Include trailing spaces
+ String[] rawValues = fields[targetCanonIndex].split("\\|");
+ List values = new ArrayList<>();
+ for (String rawValue : rawValues) {
+ values.add(simpleCanonicalize(rawValue));
+ }
+ idToValue.put(fields[exIdIndex], new ListValue(values));
+ }
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ LogInfo.end_track();
+ }
+
+ /**
+ * Like canonicalize, but assume that the string is already well-formed:
+ * - A number should look like a float
+ * - A date should be in the ISO format
+ * - Otherwise, the value is treated as a string.
+ */
+ protected Value simpleCanonicalize(String origString) {
+ Value answer;
+ // Try converting to a number.
+ answer = StringNormalizationUtils.parseNumberStrict(origString);
+ if (answer != null) return answer;
+ // Try converting to a date.
+ answer = StringNormalizationUtils.parseDate(origString);
+ if (answer != null) return answer;
+ // Just treat as a description string
+ return new DescriptionValue(origString);
+ }
+
}
diff --git a/src/edu/stanford/nlp/sempre/tables/alter/AggregatedTurkData.java b/src/edu/stanford/nlp/sempre/tables/alter/AggregatedTurkData.java
index 0ebbf3e..e6ff578 100644
--- a/src/edu/stanford/nlp/sempre/tables/alter/AggregatedTurkData.java
+++ b/src/edu/stanford/nlp/sempre/tables/alter/AggregatedTurkData.java
@@ -93,7 +93,7 @@ public class AggregatedTurkData {
List values = new ArrayList<>();
for (String x : response.split("\\|"))
values.add(new DescriptionValue(StringNormalizationUtils.unescapeTSV(x)));
- return TargetValuePreprocessor.getSingleton().preprocess(new ListValue(values));
+ return TargetValuePreprocessor.getSingleton().preprocess(new ListValue(values), null);
}
}
diff --git a/src/edu/stanford/nlp/sempre/tables/features/AnchorFeatureComputer.java b/src/edu/stanford/nlp/sempre/tables/features/AnchorFeatureComputer.java
new file mode 100644
index 0000000..0b4f098
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/tables/features/AnchorFeatureComputer.java
@@ -0,0 +1,67 @@
+package edu.stanford.nlp.sempre.tables.features;
+
+import java.util.*;
+
+import edu.stanford.nlp.sempre.*;
+import edu.stanford.nlp.sempre.tables.StringNormalizationUtils;
+import edu.stanford.nlp.sempre.tables.TableKnowledgeGraph;
+import edu.stanford.nlp.sempre.tables.TableCell;
+import edu.stanford.nlp.sempre.tables.TableColumn;
+import fig.basic.LogInfo;
+
+public class AnchorFeatureComputer implements FeatureComputer {
+
+ @Override
+ public void extractLocal(Example ex, Derivation deriv) {
+ if (!(FeatureExtractor.containsDomain("anchored-entity"))) return;
+ if (!(deriv.rule.sem instanceof FuzzyMatchFn)) return;
+ FuzzyMatchFn sem = (FuzzyMatchFn) deriv.rule.sem;
+ if (sem.getMatchAny() || sem.getMode() != FuzzyMatchFn.FuzzyMatchFnMode.ENTITY) return;
+ String phrase = ((StringValue) ((ValueFormula>) deriv.child(0).formula).value).value;
+ NameValue predicate = (NameValue) ((ValueFormula>) deriv.formula).value;
+ TableKnowledgeGraph graph = (TableKnowledgeGraph) ex.context.graph;
+ extractMatchingFeatures(graph, deriv, phrase, predicate);
+ }
+
+ private void extractMatchingFeatures(TableKnowledgeGraph graph,
+ Derivation deriv, String phrase, NameValue predicate) {
+ String predicateString = graph.getOriginalString(predicate);
+ //LogInfo.logs("%s -> %s = %s", phrase, predicate, predicateString);
+ predicateString = StringNormalizationUtils.simpleNormalize(predicateString).toLowerCase();
+ if (predicateString.equals(phrase)) {
+ deriv.addFeature("a-e", "exact");
+ //LogInfo.logs("%s %s exact", phrase, predicateString);
+ } else if (predicateString.startsWith(phrase + " ")) {
+ deriv.addFeature("a-e", "prefix");
+ //LogInfo.logs("%s %s prefix", phrase, predicateString);
+ } else if (predicateString.endsWith(" " + phrase)) {
+ deriv.addFeature("a-e", "suffix");
+ //LogInfo.logs("%s %s suffix", phrase, predicateString);
+ } else if (predicateString.contains(" " + phrase + " ")){
+ deriv.addFeature("a-e", "substring");
+ //LogInfo.logs("%s %s substring", phrase, predicateString);
+ } else {
+ deriv.addFeature("a-e", "other");
+ //LogInfo.logs("%s %s other", phrase, predicateString);
+ }
+ // Does the phrase match other cells?
+ Set matches = new HashSet<>();
+ for (TableColumn column : graph.columns) {
+ for (TableCell cell : column.children) {
+ String s = StringNormalizationUtils.simpleNormalize(cell.properties.originalString).toLowerCase();
+ if (s.contains(phrase) && !cell.properties.id.equals(predicate.id)) {
+ matches.add(s);
+ }
+ }
+ }
+ //LogInfo.logs(">> %s", matches);
+ if (matches.size() == 0) {
+ deriv.addFeature("a-e", "unique");
+ } else if (matches.size() < 3) {
+ deriv.addFeature("a-e", "multiple;" + matches.size());
+ } else {
+ deriv.addFeature("a-e", "multiple;>=3");
+ }
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/tables/features/ColumnCategoryInfo.java b/src/edu/stanford/nlp/sempre/tables/features/ColumnCategoryInfo.java
new file mode 100644
index 0000000..d776fc2
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/tables/features/ColumnCategoryInfo.java
@@ -0,0 +1,84 @@
+package edu.stanford.nlp.sempre.tables.features;
+
+import java.io.*;
+import java.util.*;
+
+import edu.stanford.nlp.sempre.Example;
+import edu.stanford.nlp.sempre.tables.TableKnowledgeGraph;
+import fig.basic.*;
+
+public class ColumnCategoryInfo {
+ public static class Options {
+ @Option(gloss = "Read category information from this file")
+ public String tableCategoryInfo = null;
+ }
+ public static Options opts = new Options();
+
+ // ============================================================
+ // Singleton access
+ // ============================================================
+
+ private static ColumnCategoryInfo singleton;
+
+ public static ColumnCategoryInfo getSingleton() {
+ if (opts.tableCategoryInfo == null)
+ return null;
+ else if (singleton == null)
+ singleton = new ColumnCategoryInfo();
+ return singleton;
+ }
+
+ // ============================================================
+ // Read data from file
+ // ============================================================
+
+ // tableId -> columnIndex -> list of (category, weight)
+ protected static Map>>> allCategoryInfo = null;
+
+ private ColumnCategoryInfo() {
+ LogInfo.begin_track("Loading category information from %s", opts.tableCategoryInfo);
+ allCategoryInfo = new HashMap<>();
+ try {
+ BufferedReader reader = IOUtils.openIn(opts.tableCategoryInfo);
+ String line;
+ while ((line = reader.readLine()) != null) {
+ String[] tokens = line.split("\t");
+ String tableId = tokens[0];
+ List>> categoryInfoForTable = allCategoryInfo.get(tableId);
+ if (categoryInfoForTable == null)
+ allCategoryInfo.put(tableId, categoryInfoForTable = new ArrayList<>());
+ int columnIndex = Integer.parseInt(tokens[1]);
+ // Assume that the columns are ordered
+ assert categoryInfoForTable.size() == columnIndex;
+ // Read the category-weight pairs
+ List> categories = new ArrayList<>();
+ for (int i = 2; i < tokens.length; i++) {
+ String[] pair = tokens[i].split(":");
+ categories.add(new Pair<>(pair[0], Double.parseDouble(pair[1])));
+ }
+ categoryInfoForTable.add(categories);
+ }
+ reader.close();
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ LogInfo.end_track();
+ }
+
+ // ============================================================
+ // Getters
+ // ============================================================
+
+ public List> get(String tableId, int columnIndex) {
+ return allCategoryInfo.get(tableId).get(columnIndex);
+ }
+
+ public List> get(Example ex, String columnId) {
+ TableKnowledgeGraph graph = (TableKnowledgeGraph) ex.context.graph;
+ String tableId = graph.filename;
+ int columnIndex = graph.getColumnIndex(columnId);
+ if (columnIndex == -1) return null;
+ return allCategoryInfo.get(tableId).get(columnIndex);
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/tables/features/PhraseDenotationFeatureComputer.java b/src/edu/stanford/nlp/sempre/tables/features/PhraseDenotationFeatureComputer.java
index 5f33941..895ef39 100644
--- a/src/edu/stanford/nlp/sempre/tables/features/PhraseDenotationFeatureComputer.java
+++ b/src/edu/stanford/nlp/sempre/tables/features/PhraseDenotationFeatureComputer.java
@@ -113,6 +113,7 @@ public class PhraseDenotationFeatureComputer implements FeatureComputer {
LogInfo.logs("%s %s %s", deriv.value, deriv.type, denotationTypes);
for (String denotationType : denotationTypes) {
for (PhraseInfo phraseInfo : phraseInfos) {
+ if (PhraseInfo.opts.forbidBorderStopWordInLexicalizedFeatures && phraseInfo.isBorderStopWord) continue;
deriv.addFeature("p-d", phraseInfo.lemmaText + ";" + denotationType);
}
// Check original column text
diff --git a/src/edu/stanford/nlp/sempre/tables/features/PhraseInfo.java b/src/edu/stanford/nlp/sempre/tables/features/PhraseInfo.java
index 1ba0cd5..86f06fc 100644
--- a/src/edu/stanford/nlp/sempre/tables/features/PhraseInfo.java
+++ b/src/edu/stanford/nlp/sempre/tables/features/PhraseInfo.java
@@ -2,6 +2,7 @@ package edu.stanford.nlp.sempre.tables.features;
import java.util.*;
import java.util.concurrent.ExecutionException;
+import java.util.regex.Pattern;
import com.google.common.cache.*;
@@ -23,6 +24,8 @@ public class PhraseInfo {
public int maxPhraseLength = 3;
@Option(gloss = "Fuzzy match predicates")
public boolean computeFuzzyMatchPredicates = false;
+ @Option(gloss = "Do not produce lexicalized features if the phrase begins or ends with a stop word")
+ public boolean forbidBorderStopWordInLexicalizedFeatures = true;
}
public static Options opts = new Options();
@@ -35,6 +38,7 @@ public class PhraseInfo {
public final List nerTags;
public final String canonicalPosSeq;
public final List fuzzyMatchedPredicates;
+ public final boolean isBorderStopWord; // true if the first or last word is a stop word
public PhraseInfo(Example ex, int start, int end) {
this.start = start;
@@ -49,6 +53,7 @@ public class PhraseInfo {
lemmaText = languageInfo.lemmaPhrase(start, end).toLowerCase();
canonicalPosSeq = languageInfo.canonicalPosSeq(start, end);
fuzzyMatchedPredicates = opts.computeFuzzyMatchPredicates ? getFuzzyMatchedPredicates(ex.context) : null;
+ isBorderStopWord = isStopWord(languageInfo.lemmaTokens.get(start)) || isStopWord(languageInfo.lemmaTokens.get(end - 1));
}
private List getFuzzyMatchedPredicates(ContextValue context) {
@@ -71,6 +76,17 @@ public class PhraseInfo {
return matchedPredicates;
}
+ static final Pattern ALL_PUNCT = Pattern.compile("^[^A-Za-z0-9]*$");
+ static final Set STOP_WORDS = new HashSet<>(Arrays.asList(
+ "a", "an", "the", "be", "of", "in", "on", "do"
+ ));
+
+ static boolean isStopWord(String x) {
+ if (ALL_PUNCT.matcher(x).matches()) return true;
+ if (STOP_WORDS.contains(x)) return true;
+ return false;
+ }
+
@Override
public String toString() {
return "\"" + text + "\"";
diff --git a/src/edu/stanford/nlp/sempre/tables/features/PhrasePredicateFeatureComputer.java b/src/edu/stanford/nlp/sempre/tables/features/PhrasePredicateFeatureComputer.java
index d7b5f71..e61c2d2 100644
--- a/src/edu/stanford/nlp/sempre/tables/features/PhrasePredicateFeatureComputer.java
+++ b/src/edu/stanford/nlp/sempre/tables/features/PhrasePredicateFeatureComputer.java
@@ -35,6 +35,11 @@ public class PhrasePredicateFeatureComputer implements FeatureComputer {
public boolean lexicalizedPhrasePredicate = true;
@Option(gloss = "Maximum ngram length for lexicalize all pair features")
public int maxNforLexicalizeAllPairs = Integer.MAX_VALUE;
+ @Option(gloss = "phrase-category: Weight threshold")
+ public double phraseCategoryWeightThreshold = 0.8;
+ @Option(gloss = "phrase-category: Use binary features instead of continuous ones")
+ public boolean phraseCategoryBinary = true;
+
}
public static Options opts = new Options();
@@ -47,7 +52,8 @@ public class PhrasePredicateFeatureComputer implements FeatureComputer {
@Override
public void extractLocal(Example ex, Derivation deriv) {
if (!(FeatureExtractor.containsDomain("phrase-predicate")
- || FeatureExtractor.containsDomain("missing-predicate"))) return;
+ || FeatureExtractor.containsDomain("phrase-formula")
+ || FeatureExtractor.containsDomain("phrase-category"))) return;
// Only compute features at the root, except when the partial option is set.
if (!opts.defineOnPartialDerivs && !deriv.isRoot(ex.numTokens())) return;
List phraseInfos = PhraseInfo.getPhraseInfos(ex);
@@ -58,7 +64,8 @@ public class PhrasePredicateFeatureComputer implements FeatureComputer {
LogInfo.logs("Derivation: %s", deriv);
LogInfo.logs("Predicates: %s", predicateInfos);
}
- if (FeatureExtractor.containsDomain("phrase-predicate")) {
+ if (FeatureExtractor.containsDomain("phrase-predicate")
+ || FeatureExtractor.containsDomain("phrase-category")) {
if (opts.defineOnPartialDerivs) {
deriv.getTempState().put("p-p", new ArrayList<>(predicateInfos));
// Subtract predicates from children
@@ -109,7 +116,7 @@ public class PhrasePredicateFeatureComputer implements FeatureComputer {
private void extractMatch(Example ex, Derivation deriv,
PhraseInfo phraseInfo, String phraseString, String phraseType,
PredicateInfo predicateInfo, String predicateString, String predicateType, double factor) {
- if (opts.unlexicalizedPhrasePredicate) {
+ if (FeatureExtractor.containsDomain("phrase-predicate") && opts.unlexicalizedPhrasePredicate) {
if (phraseString.equals(predicateString)) {
defineFeatures(ex, deriv, phraseInfo, predicateInfo, phraseType + "=" + predicateType,
phraseString, predicateString, factor);
@@ -132,10 +139,27 @@ public class PhrasePredicateFeatureComputer implements FeatureComputer {
}
}
}
- if (opts.lexicalizedPhrasePredicate && phraseInfo.end - phraseInfo.start <= maxNforLexicalizeAllPairs) {
+ if (FeatureExtractor.containsDomain("phrase-predicate") && opts.lexicalizedPhrasePredicate
+ && phraseInfo.end - phraseInfo.start <= maxNforLexicalizeAllPairs
+ && (!PhraseInfo.opts.forbidBorderStopWordInLexicalizedFeatures || !phraseInfo.isBorderStopWord)) {
deriv.addFeature("p-p",
phraseType + phraseString + ";" + predicateType + predicateString, factor);
}
+ if (FeatureExtractor.containsDomain("phrase-category") && predicateInfo.type == PredicateType.BINARY
+ && (!PhraseInfo.opts.forbidBorderStopWordInLexicalizedFeatures || !phraseInfo.isBorderStopWord)) {
+ ColumnCategoryInfo catInfo = ColumnCategoryInfo.getSingleton();
+ List> categories = catInfo.get(ex, predicateInfo.predicate);
+ if (categories != null) {
+ for (Pair pair : categories) {
+ if (pair.getSecond() >= opts.phraseCategoryWeightThreshold) {
+ if (opts.phraseCategoryBinary)
+ deriv.addFeature("p-c", phraseType + phraseString + ";" + pair.getFirst());
+ else
+ deriv.addFeature("p-c", phraseType + phraseString + ";" + pair.getFirst(), pair.getSecond());
+ }
+ }
+ }
+ }
}
private void defineFeatures(Example ex, Derivation deriv, PhraseInfo phraseInfo, PredicateInfo predicateInfo,
diff --git a/src/edu/stanford/nlp/sempre/tables/lambdadcs/DenotationUtils.java b/src/edu/stanford/nlp/sempre/tables/lambdadcs/DenotationUtils.java
index afa098a..f586b3b 100644
--- a/src/edu/stanford/nlp/sempre/tables/lambdadcs/DenotationUtils.java
+++ b/src/edu/stanford/nlp/sempre/tables/lambdadcs/DenotationUtils.java
@@ -238,6 +238,8 @@ public final class DenotationUtils {
*/
public static UnaryDenotation superlativeUnary(int rank, int count, List> pairs,
SuperlativeFormula.Mode mode, TypeProcessor processor) {
+ if (rank <= 0 || count <= 0 || rank >= 1000000 || count >= 100000)
+ LogInfo.fails("Invalid superlative (rank = %d, count = %d)", rank, count);
if (pairs.isEmpty()) {
if (LambdaDCSExecutor.opts.superlativesFailOnEmptyLists)
throw new LambdaDCSException(Type.emptyList, "Cannot call %s on an empty list.", mode);
@@ -343,7 +345,7 @@ public final class DenotationUtils {
public boolean isCompatible(Value v) {
return v instanceof NumberValue;
}
-
+
@Override
public boolean isSortable(Collection values) {
return true;
@@ -387,19 +389,19 @@ public final class DenotationUtils {
public boolean isCompatible(Value v) {
return v instanceof DateValue;
}
-
+
@Override
public boolean isSortable(Collection values) {
DateValue firstDate = null;
for (Value value : values) {
- DateValue date = (DateValue) value;
+ DateValue date = (DateValue) value;
if (firstDate == null) {
firstDate = date;
} else {
if ((firstDate.year == -1) != (date.year == -1)) return false;
if ((firstDate.month == -1) != (date.month == -1)) return false;
if ((firstDate.day == -1) != (date.day == -1)) return false;
- }
+ }
}
return true;
}
@@ -465,5 +467,5 @@ public final class DenotationUtils {
throw new LambdaDCSException(Type.typeMismatch, "Cannot compare values");
}
}
-
+
}
diff --git a/src/edu/stanford/nlp/sempre/tables/lambdadcs/LambdaDCSExecutor.java b/src/edu/stanford/nlp/sempre/tables/lambdadcs/LambdaDCSExecutor.java
index 39f85d7..0d468b2 100644
--- a/src/edu/stanford/nlp/sempre/tables/lambdadcs/LambdaDCSExecutor.java
+++ b/src/edu/stanford/nlp/sempre/tables/lambdadcs/LambdaDCSExecutor.java
@@ -292,6 +292,9 @@ class LambdaDCSCoreLogic {
computeUnary(superlative.rank, typeHint.unrestrictedUnary()).range());
int count = DenotationUtils.getSinglePositiveInteger(
computeUnary(superlative.count, typeHint.unrestrictedUnary()).range());
+ if (rank != 1 || count != 1) {
+ LogInfo.logs("Superlative WTF: %s | rank %d | count %d", formula, rank, count);
+ }
Unarylike headD = computeUnary(superlative.head, typeHint);
Binarylike relationD;
if (superlative.relation instanceof ReverseFormula) {
diff --git a/src/edu/stanford/nlp/sempre/tables/test/BatchTableExecutor.java b/src/edu/stanford/nlp/sempre/tables/test/BatchTableExecutor.java
new file mode 100644
index 0000000..1dab7f6
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/tables/test/BatchTableExecutor.java
@@ -0,0 +1,136 @@
+package edu.stanford.nlp.sempre.tables.test;
+
+import java.io.*;
+import java.util.*;
+
+import edu.stanford.nlp.sempre.*;
+import edu.stanford.nlp.sempre.tables.TableKnowledgeGraph;
+import edu.stanford.nlp.sempre.tables.TableValueEvaluator;
+import edu.stanford.nlp.sempre.tables.lambdadcs.LambdaDCSExecutor;
+import fig.basic.*;
+import fig.exec.Execution;
+
+/**
+ * Execute the specified logical forms on the specified WikiTableQuestions context.
+ *
+ * @author ppasupat
+ */
+public class BatchTableExecutor implements Runnable {
+ public static class Options {
+ @Option(gloss = "TSV file containing table contexts and logical forms")
+ public String batchInput;
+ @Option(gloss = "Datasets for mapping example IDs to contexts")
+ public List batchDatasets = Arrays.asList("lib/data/tables/data/training.examples");
+ }
+ public static Options opts = new Options();
+
+ public static void main(String[] args) {
+ Execution.run(args, "BatchTableExecutorMain", new BatchTableExecutor(), Master.getOptionsParser());
+ }
+
+ @Override
+ public void run() {
+ if (opts.batchInput == null || opts.batchInput.isEmpty()) {
+ LogInfo.logs("*******************************************************************************");
+ LogInfo.logs("USAGE: ./run @mode=tables @class=execute -batchInput ");
+ LogInfo.logs("");
+ LogInfo.logs("Input file format: Each line has something like");
+ LogInfo.logs(" nt-218 [tab] (count (fb:type.object.type fb:type.row))");
+ LogInfo.logs("or");
+ LogInfo.logs(" csv/204-csv/23.csv [tab] (count (fb:type.object.type fb:type.row))");
+ LogInfo.logs("");
+ LogInfo.logs("Results will also be printed to state/execs/___.exec/denotations.tsv");
+ LogInfo.logs("Output format:");
+ LogInfo.logs(" nt-218 [tab] (count (fb:type.object.type fb:type.row)) [tab] (list (number 10)) [tab] false");
+ LogInfo.logs("where the last column indicates whether the answer is consistent with the target answer");
+ LogInfo.logs("(only available when the first column is nt-___)");
+ LogInfo.logs("*******************************************************************************");
+ System.exit(1);
+ }
+ LambdaDCSExecutor executor = new LambdaDCSExecutor();
+ ValueEvaluator evaluator = new TableValueEvaluator();
+ try {
+ BufferedReader reader = IOUtils.openIn(opts.batchInput);
+ PrintWriter output = IOUtils.openOut(Execution.getFile("denotations.tsv"));
+ String line;
+ while ((line = reader.readLine()) != null) {
+ String[] tokens = line.split("\t");
+ String answer;
+ try {
+ Formula formula = Formula.fromString(tokens[1]);
+ if (tokens[0].startsWith("csv")) {
+ TableKnowledgeGraph graph = TableKnowledgeGraph.fromFilename(tokens[0]);
+ ContextValue context = new ContextValue(graph);
+ Value denotation = executor.execute(formula, context).value;
+ if (denotation instanceof ListValue)
+ denotation = addOriginalStrings((ListValue) denotation, graph);
+ answer = denotation.toString();
+ } else {
+ Example ex = exIdToExample(tokens[0]);
+ Value denotation = executor.execute(formula, ex.context).value;
+ if (denotation instanceof ListValue)
+ denotation = addOriginalStrings((ListValue) denotation, (TableKnowledgeGraph) ex.context.graph);
+ answer = denotation.toString();
+ boolean correct = evaluator.getCompatibility(ex.targetValue, denotation) == 1.;
+ answer = denotation.toString() + "\t" + correct;
+ }
+ } catch (Exception e) {
+ answer = "ERROR: " + e;
+ }
+ System.out.printf("%s\t%s\t%s\n", tokens[0], tokens[1], answer);
+ output.printf("%s\t%s\t%s\n", tokens[0], tokens[1], answer);
+ }
+ reader.close();
+ output.close();
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ }
+
+ private Map exIdToExampleMap;
+
+ private Example exIdToExample(String exId) {
+ if (exIdToExampleMap == null) {
+ exIdToExampleMap = new HashMap<>();
+ try {
+ for (String filename : opts.batchDatasets) {
+ BufferedReader reader = IOUtils.openIn(filename);
+ String line;
+ while ((line = reader.readLine()) != null) {
+ LispTree tree = LispTree.proto.parseFromString(line);
+ if (!"id".equals(tree.child(1).child(0).value))
+ throw new RuntimeException("Malformed example: " + line);
+ exIdToExampleMap.put(tree.child(1).child(1).value, tree);
+ }
+ }
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ }
+ Object obj = exIdToExampleMap.get(exId);
+ if (obj == null) return null;
+ Example ex;
+ if (obj instanceof LispTree) {
+ ex = Example.fromLispTree((LispTree) obj, exId);
+ ex.preprocess();
+ exIdToExampleMap.put(exId, ex);
+ } else {
+ ex = (Example) obj;
+ }
+ return ex;
+ }
+
+ ListValue addOriginalStrings(ListValue answers, TableKnowledgeGraph graph) {
+ List values = new ArrayList<>();
+ for (Value value : answers.values) {
+ if (value instanceof NameValue) {
+ NameValue name = (NameValue) value;
+ if (name.description == null)
+ value = new NameValue(name.id, graph.getOriginalString(((NameValue) value).id));
+ }
+ values.add(value);
+ }
+ return new ListValue(values);
+ }
+
+}
diff --git a/src/edu/stanford/nlp/sempre/tables/test/TableColumnAnalyzer.java b/src/edu/stanford/nlp/sempre/tables/test/TableColumnAnalyzer.java
new file mode 100644
index 0000000..a777b37
--- /dev/null
+++ b/src/edu/stanford/nlp/sempre/tables/test/TableColumnAnalyzer.java
@@ -0,0 +1,237 @@
+package edu.stanford.nlp.sempre.tables.test;
+
+import java.io.*;
+import java.util.*;
+import java.util.regex.*;
+
+import edu.stanford.nlp.sempre.*;
+import edu.stanford.nlp.sempre.tables.*;
+import fig.basic.*;
+import fig.exec.Execution;
+
+/**
+ * Analyze table columns and print out any hard-to-process column.
+ *
+ * @author ppasupat
+ */
+public class TableColumnAnalyzer implements Runnable {
+ public static class Options {
+ @Option(gloss = "Maximum number of tables to process (for debugging)")
+ public int maxNumTables = Integer.MAX_VALUE;
+ @Option(gloss = "Load Wikipedia article titles from this file")
+ public String wikiTitles = null;
+ }
+ public static Options opts = new Options();
+
+ public static void main(String[] args) {
+ Execution.run(args, "TableColumnAnalyzerMain", new TableColumnAnalyzer(), Master.getOptionsParser());
+ }
+
+ PrintWriter out;
+ PrintWriter outCompact;
+
+ @Override
+ public void run() {
+ out = IOUtils.openOutHard(Execution.getFile("column-stats.tsv"));
+ outCompact = IOUtils.openOutHard(Execution.getFile("column-compact.tsv"));
+ Map> tableIdToExIds = getTableIds();
+ int tablesProcessed = 0;
+ for (Map.Entry> entry : tableIdToExIds.entrySet()) {
+ Execution.putOutput("example", tablesProcessed);
+ String tableId = entry.getKey(),
+ tableIdAbbrev = tableId.replaceAll("csv/(\\d+)-csv/(\\d+)\\.csv", "$1-$2");
+ LogInfo.begin_track("Processing %s ...", tableId);
+ TableKnowledgeGraph graph = TableKnowledgeGraph.fromFilename(tableId);
+ out.printf("%s\tIDS\t%s\n", tableIdAbbrev, String.join(" ", entry.getValue()));
+ out.printf("%s\tCOLUMNS\t%d\n", tableIdAbbrev, graph.numColumns());
+ for (int i = 0; i < graph.numColumns(); i++) {
+ analyzeColumn(graph, graph.columns.get(i), tableIdAbbrev + "\t" + i);
+ }
+ LogInfo.end_track();
+ if (tablesProcessed++ >= opts.maxNumTables) break;
+ }
+ out.close();
+ outCompact.close();
+ }
+
+ protected Map> getTableIds() {
+ Map> tableIdToExIds = new LinkedHashMap<>();
+ LogInfo.begin_track_printAll("Collect table IDs");
+ for (Pair pathPair : Dataset.opts.inPaths) {
+ String group = pathPair.getFirst();
+ String path = pathPair.getSecond();
+ Execution.putOutput("group", group);
+ LogInfo.begin_track("Reading %s", path);
+ Iterator trees = LispTree.proto.parseFromFile(path);
+ while (trees.hasNext()) {
+ LispTree tree = trees.next();
+ if ("metadata".equals(tree.child(0).value)) continue;
+ String exId = null, tableId = null;
+ for (int i = 1; i < tree.children.size(); i++) {
+ LispTree arg = tree.child(i);
+ String label = arg.child(0).value;
+ if ("id".equals(label)) {
+ exId = arg.child(1).value;
+ } else if ("context".equals(label)) {
+ tableId = arg.child(1).child(2).value;
+ }
+ }
+ if (exId != null && tableId != null) {
+ List exIdsForTable = tableIdToExIds.get(tableId);
+ if (exIdsForTable == null)
+ tableIdToExIds.put(tableId, exIdsForTable = new ArrayList<>());
+ exIdsForTable.add(exId);
+ }
+ }
+ LogInfo.end_track();
+ }
+ LogInfo.end_track();
+ LogInfo.logs("Got %d IDs", tableIdToExIds.size());
+ return tableIdToExIds;
+ }
+
+ protected void analyzeColumn(TableKnowledgeGraph graph, TableColumn column, String printPrefix) {
+ List escapedCells = new ArrayList<>();
+ // Print the header
+ String h = column.originalString, escapedH = StringNormalizationUtils.escapeTSV(h);
+ out.printf("%s\t0\t%s\n", printPrefix, escapedH);
+ escapedCells.add(escapedH);
+ // Print the cells
+ Map typeCounts = new HashMap<>();
+ for (int j = 0; j < column.children.size(); j++) {
+ TableCell cell = column.children.get(j);
+ String c = cell.properties.originalString, escapedC = StringNormalizationUtils.escapeTSV(c);
+ escapedCells.add(escapedC);
+ // Infer the type
+ List types = analyzeCell(c);
+ for (String type : types)
+ MapUtils.incr(typeCounts, type);
+ out.printf("%s\t%d\t%s\t%s\n", printPrefix, j + 1, String.join("|", types), escapedC);
+ }
+ // Analyze the common types
+ List commonTypes = new ArrayList<>();
+ for (Map.Entry entry : typeCounts.entrySet()) {
+ if (entry.getValue() == column.children.size()) {
+ commonTypes.add(entry.getKey());
+ } else if (entry.getValue() == column.children.size() - 1) {
+ commonTypes.add("ALMOST-" + entry.getKey());
+ }
+ }
+ outCompact.printf("%s\t%s\t%s\n", String.join("|", commonTypes), printPrefix, String.join("\t", escapedCells));
+ }
+
+ // ============================================================
+ // Cell analysis
+ // ============================================================
+
+ public static final Pattern ORDINAL = Pattern.compile("^(\\d+)(st|nd|rd|th)$");
+
+ protected List analyzeCell(String c) {
+ List types = new ArrayList<>();
+ LanguageInfo languageInfo = LanguageAnalyzer.getSingleton().analyze(c);
+ {
+ // Integer
+ NumberValue n = StringNormalizationUtils.parseNumberStrict(c);
+ if (n != null) {
+ // Number
+ types.add("num");
+ // Integer
+ double value = n.value;
+ if (Math.abs(value - Math.round(value)) < 1e-9) {
+ types.add("int");
+ if (c.matches("^[12]\\d\\d\\d$")) {
+ // Year?
+ types.add("year");
+ }
+ }
+ }
+ }
+ {
+ // Ordinal
+ Matcher m = ORDINAL.matcher(c);
+ if (m.matches()) {
+ types.add("ordinal");
+ }
+ }
+ {
+ // Integer-Integer
+ String[] splitted = StringNormalizationUtils.STRICT_DASH.split(c);
+ if (splitted.length == 2 && splitted[0].matches("^[0-9]+$") && splitted[1].matches("^[0-9]+$")) {
+ types.add("2ints");
+ }
+ }
+ {
+ // Date
+ DateValue date = StringNormalizationUtils.parseDateWithLanguageAnalyzer(languageInfo);
+ if (date != null) {
+ types.add("date");
+ // Also more detailed date type
+ types.add("date-"
+ + (date.year != -1 ? "Y" : "")
+ + (date.month != -1 ? "M" : "")
+ + (date.day != -1 ? "D" : ""));
+ }
+ }
+ {
+ // Quoted text
+ if (c.matches("^[“”\"].*[“”\"]$")) {
+ types.add("quoted");
+ }
+ }
+ if (opts.wikiTitles != null) {
+ // Wikipedia titles
+ WikipediaTitleLibrary library = WikipediaTitleLibrary.getSingleton();
+ if (library.contains(c)) {
+ types.add("wiki");
+ }
+ }
+ {
+ // POS and NER
+ types.add("POS=" + String.join("-", languageInfo.posTags));
+ types.add("NER=" + String.join("-", languageInfo.nerTags));
+ }
+ return types;
+ }
+
+ // ============================================================
+ // Helper class: Wikipedia titles
+ // ============================================================
+
+ public static class WikipediaTitleLibrary {
+
+ private static WikipediaTitleLibrary _singleton = null;
+
+ public static WikipediaTitleLibrary getSingleton() {
+ if (_singleton == null)
+ _singleton = new WikipediaTitleLibrary();
+ return _singleton;
+ }
+
+ Set titles = new HashSet<>();
+
+ private WikipediaTitleLibrary() {
+ assert opts.wikiTitles != null;
+ LogInfo.begin_track("Reading Wikipedia article titles from %s ...", opts.wikiTitles);
+ try {
+ BufferedReader reader = IOUtils.openIn(opts.wikiTitles);
+ String line;
+ while ((line = reader.readLine()) != null) {
+ titles.add(line);
+ if (titles.size() <= 10) {
+ LogInfo.logs("Example title: %s", line);
+ }
+ }
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ LogInfo.logs("Read %d titles", titles.size());
+ LogInfo.end_track();
+ }
+
+ public boolean contains(String c) {
+ return titles.contains(c.toLowerCase().trim());
+ }
+ }
+
+
+}
diff --git a/tables/README.md b/tables/README.md
index 2197eb8..cc1c98e 100644
--- a/tables/README.md
+++ b/tables/README.md
@@ -33,7 +33,7 @@ Running the code
./run @mode=tables @data=u-1 @feat=all @train=1 -maxex train,100 dev,100
- The command should take less than an hour.
+ The command should take less than 30 minutes.
* To train on the complete development set, remove `-maxex train,100 dev,100`
@@ -41,6 +41,25 @@ Running the code
Other available sets include `u-2`, ..., `u-5` (four other development splits)
and `test` (actual train-test split).
+Other options
+-------------
+
+### Macro Grammar (Experimental)
+
+Macro grammar can be used to significantly speed up the parser.
+To turn on macro grammar, run the following:
+
+ ./run @mode=tables @data=u-1 @feat=more @parser=cprune @grammar=extended @fuzzy=editdist-fuzzy @train=1
+
+Please refer to the following paper for more information:
+
+> Yuchen Zhang, Panupong Pasupat, Percy Liang.
+> Macro Grammars and Holistic Triggering for Efficient Semantic Parsing
+> Empirical Methods on Natural Language Processing (EMNLP), 2017.
+
+Currently the module does not support model saving, and testing has to be done on the official test set.
+These features will be added in the future.
+
Official evaluation
-------------------
diff --git a/tables/grammars/extended.grammar b/tables/grammars/extended.grammar
new file mode 100644
index 0000000..d45e1b5
--- /dev/null
+++ b/tables/grammars/extended.grammar
@@ -0,0 +1,169 @@
+# (Extended) Generic Grammar
+# Use more generic compositional patterns.
+
+################################################################
+# Macros
+
+(def @R reverse)
+(def @type fb:type.object.type)
+(def @row fb:type.row)
+
+(def @next fb:row.row.next)
+(def @!next !fb:row.row.next)
+(def @index fb:row.row.index)
+(def @!index !fb:row.row.index)
+
+(def @p.num fb:cell.cell.number)
+(def @!p.num !fb:cell.cell.number)
+(def @p.date fb:cell.cell.date)
+(def @!p.date !fb:cell.cell.date)
+(def @p.second fb:cell.cell.second)
+(def @!p.second !fb:cell.cell.second)
+
+################################################################
+# Lexicon
+
+################################
+# Anchored Rules: Entity, Unary, Binary
+(rule $Entity ($PHRASE) (FuzzyMatchFn entity) (anchored 1))
+#(rule $Binary ($PHRASE) (FuzzyMatchFn binary) (anchored 1))
+(rule $Entity ($PHRASE) (NumberFn) (anchored 1))
+(rule $Entity ($PHRASE) (DateFn) (anchored 1))
+
+################################
+# Create binary from thin air
+(rule $Binary (nothing) (FuzzyMatchFn any binary))
+(rule $Unary (nothing) (FuzzyMatchFn any unary))
+
+################################
+# Property
+(for @property (@p.num @p.date)
+ (rule $Property (nothing) (ConstantFn @property))
+)
+(when second
+ (rule $Property (nothing) (ConstantFn @p.second))
+)
+
+################################
+# Generic RowSet
+(rule $RowSet (nothing) (ConstantFn (@type @row)))
+
+################################
+# Anchored ValueSet
+(rule $ValueSet ($Entity) (IdentityFn))
+
+# [TAG] alternative: "X or Y" questions
+(when alternative
+ (rule $ValueSet ($Entity $Entity)
+ (lambda e1 (lambda e2 (or (var e1) (var e2))))
+ )
+)
+
+################################
+# Join
+
+(rule $RowSet ($Binary $ValueSet) (lambda b (lambda v ((var b) (var v)))))
+
+(rule $ValueSet ($Binary $RowSet) (lambda b (lambda r ((@R (var b)) (var r)))))
+
+(rule $RowSet ($Binary $Property $ValueSet)
+ (lambda b (lambda p (lambda v ((var b) ((var p) (var v))))))
+)
+
+(rule $ValueSet ($Binary $Property $RowSet)
+ (lambda b (lambda p (lambda r ((@R (var p)) ((@R (var b)) (var r))))))
+)
+
+# [TAG] movement: "next" / "previous"
+(when movement
+ (for @movement (@next @!next)
+ (rule $RowSet ($RowSet) (lambda r (@movement (var r))))
+ )
+)
+
+# [TAG] comparison: "at least" / "more than"
+(when comparison
+ (for @comparison (< > <= >=)
+ (rule $RowSet ($Binary $Property $Entity)
+ (lambda b (lambda p (lambda e ((var b) ((var p) (@comparison (var e)))))))
+ )
+ )
+)
+
+# [TAG] != : "not zero" / "same"
+(when neq
+ (rule $RowSet ($Binary $Entity) (lambda b (lambda e ((var b) (!= (var e))))))
+ (rule $RowSet ($Binary $Property $Entity)
+ (lambda b (lambda p (lambda e ((var b) ((var p) (!= (var e)))))))
+ )
+)
+
+################################
+# Aggregate
+
+(when count
+ (rule $SingleValue ($RowSet) (lambda r (count (var r))))
+)
+
+(when aggregate
+ (rule $SingleValue ($ValueSet) (lambda r (min (var r))))
+ (rule $SingleValue ($ValueSet) (lambda r (max (var r))))
+ (rule $SingleValue ($ValueSet) (lambda r (sum (var r))))
+ (rule $SingleValue ($ValueSet) (lambda r (avg (var r))))
+)
+
+################################
+# Superlative
+
+(rule $FnOnRow ($Binary $Property)
+ (lambda b (lambda p (lambda x ((@R (var p)) ((@R (var b)) (var x))))))
+)
+
+(when superlative
+ (for @argm (argmax argmin)
+ (rule $RowSet ($RowSet) (lambda r (@argm 1 1 (var r) @index)))
+ (rule $RowSet ($RowSet $FnOnRow) (lambda r (lambda f (@argm 1 1 (var r) (@R (var f))))))
+ )
+)
+
+################################
+# Merge
+
+(when merge
+ (rule $RowSet ($RowSet $RowSet)
+ (lambda r1 (lambda r2 (and (var r1) (var r2))))
+ )
+)
+
+################################
+# Arithmatic
+
+(rule $FnOnValue ($Binary $Binary $Property)
+ (lambda b1 (lambda b2 (lambda p (lambda x ((@R (var p)) ((@R (var b2)) ((var b1) (var x))))))))
+)
+(rule $FnOnValue ($Binary)
+ (lambda b (lambda x (count ((var b) (var x)))))
+)
+(rule $FnOnValue ($Binary $Property)
+ (lambda b (lambda p (lambda x (count ((var b) ((var p) (var x)))))))
+)
+
+(when arithmetic
+ (rule $SingleValue ($FnOnValue $Entity $Entity)
+ (lambda f (lambda e1 (lambda e2 (- ((var f) (var e1)) ((var f) (var e2))))))
+ )
+)
+
+################################
+# V-superlative
+
+(when v-superlative
+ (for @argm (argmax argmin)
+ (rule $ValueSet ($ValueSet $FnOnValue) (lambda v (lambda f (@argm 1 1 (var v) (reverse (var f))))))
+ )
+)
+
+################################
+# ROOT
+(rule $ROOT ($ValueSet) (IdentityFn))
+(rule $ROOT ($SingleValue) (IdentityFn))
diff --git a/tables/wikipedia-scripts/download-wikipedia-pages.py b/tables/wikipedia-scripts/download-wikipedia-pages.py
index 516a2b4..0d3e4eb 100755
--- a/tables/wikipedia-scripts/download-wikipedia-pages.py
+++ b/tables/wikipedia-scripts/download-wikipedia-pages.py
@@ -28,7 +28,7 @@ def download(data, i, outdir):
hashcode = data['hashcode'] = CACHE.get_hashcode(url)
data['url'] = url
result = CACHE.get_page(url)
- if result is not None:
+ if result is not None and not os.path.exists(os.path.join(outdir, str(data['id']) + '.html')):
os.symlink(os.path.join('..', 'web.cache', hashcode),
os.path.join(outdir, str(data['id']) + '.html'))
with open(os.path.join(outdir, str(data['id']) + '.json'), 'w') as fout:
diff --git a/tables/wikipedia-scripts/find-good-tables.py b/tables/wikipedia-scripts/find-good-tables.py
index 4f9d29a..8cf1c1c 100755
--- a/tables/wikipedia-scripts/find-good-tables.py
+++ b/tables/wikipedia-scripts/find-good-tables.py
@@ -49,10 +49,10 @@ class TableStat(object):
self.num_min = min(table.num_rows, table.num_cols)
self.num_max = max(table.num_rows, table.num_cols)
self.num_cells = table.num_cells
- self.num_empty_cells = len([x for x in table.cells if not x])
- self.num_long = len([x for x in table.cells if len(x) >= 40])
+ self.num_empty_cells = len([x for x in table.cells if not x[1]])
+ self.num_long = len([x for x in table.cells if len(x[1]) >= 40])
self.num_short_headers = len([x for x in table.rows[0] if len(x) <= 3])
- self.num_numeric_cells = len([x for x in table.cells if re.search('[0-9]', x)])
+ self.num_numeric_cells = len([x for x in table.cells if re.search('[0-9]', x[1])])
self.num_repetitive_cols = sum([self.is_repetitive(col) for col in table.cols], 0)
self.num_similar_colpairs = sum([self.are_similar(col1, col2)
for (col1, col2) in combinations(table.cols, 2)], 0)
@@ -61,11 +61,11 @@ class TableStat(object):
self.get_scores()
def is_repetitive(self, col):
- return len(set(get_repeatable_token(x) for x in col)) < 0.5 * len(col)
+ return len(set(get_repeatable_token(x[1]) for x in col)) < 0.5 * len(col)
def are_similar(self, col1, col2):
- col1 = set(get_alpha(x) for x in col1) - set([''])
- col2 = set(get_alpha(x) for x in col2) - set([''])
+ col1 = set(get_alpha(x[1]) for x in col1) - set([''])
+ col2 = set(get_alpha(x[1]) for x in col2) - set([''])
return len(col1 & col2) >= 3
def __str__(self):
@@ -119,6 +119,39 @@ def check_table(table, criterion):
and stat.num_short_headers < stat.num_cols * 0.3):
return stat
+def dump(i, page_id, table_id, stat, args):
+ with open(os.path.join(args.source_dir, '%d.json' % page_id), 'r', 'utf8') as fin:
+ meta = json.load(fin)
+ meta['tableIndex'] = table_id
+ if args.custom_filenames:
+ while True:
+ try:
+ i = int(raw_input('filename for "%s": ' % meta['title']))
+ break
+ except:
+ pass
+ with open(os.path.join(args.json_dir, '%d.json' % i), 'w', 'utf8') as fout:
+ json.dump(meta, fout)
+ if args.copy:
+ shutil.copy(os.path.join(args.source_dir, '%d.html' % page_id),
+ os.path.join(args.page_dir, '%d.html' % i))
+ else:
+ os.symlink(os.path.relpath(os.path.join(args.source_dir, '%d.html' % page_id),
+ args.page_dir),
+ os.path.join(args.page_dir, '%d.html' % i))
+ print '>' * 30, 'Written', i, page_id, table_id, '<' * 30
+ print stat
+
+def read_used_pages(json_files, args):
+ used_pages = set()
+ for filename in json_files:
+ with open(os.path.join(args.json_dir, filename)) as fin:
+ data = json.load(fin)
+ used_pages.add(data['id'])
+ print 'Found {} used pages'.format(len(used_pages))
+ return used_pages
+
+
def main():
parser = argparse.ArgumentParser()
parser.add_argument('-n', '--amount', type=int, default=20,
@@ -130,26 +163,41 @@ def main():
help="criterion for finding tables (higher = more strict)")
parser.add_argument('-d', '--dry-run', action='store_true',
help="do not copy files")
+ parser.add_argument('-C', '--copy', action='store_true',
+ help='copy the html file instead of making a symlink')
parser.add_argument('-i', '--fixed-id', type=int,
help="use this fixed page id")
parser.add_argument('--custom-filenames', action='store_true',
help="set custom filename for each file")
+ parser.add_argument('-q', '--quick', action='store_true',
+ help='dump the good tables immediately after they are found')
+ parser.add_argument('-r', '--resume-quick', action='store_true',
+ help='continue the quick mode')
parser.add_argument('outprefix')
args = parser.parse_args()
- json_dir = os.path.join(args.outprefix + '-json')
- page_dir = os.path.join(args.outprefix + '-page')
- if os.path.exists(json_dir) or os.path.exists(page_dir):
+ quick_index, quick_used_pages = 0, set()
+ args.json_dir = os.path.join(args.outprefix + '-json')
+ args.page_dir = os.path.join(args.outprefix + '-page')
+ if args.resume_quick:
+ # Read the highest number in json_dir
+ json_files = os.listdir(args.json_dir)
+ print 'Found {} JSON files'.format(len(json_files))
+ quick_index = max(int(x.replace('.json', '')) for x in json_files) + 1
+ quick_used_pages = read_used_pages(json_files, args)
+ elif os.path.exists(args.json_dir) or os.path.exists(args.page_dir):
if raw_input('Path exists. Clobber? ')[0:].lower() != 'y':
exit(1)
elif not args.dry_run:
- os.makedirs(json_dir)
- os.makedirs(page_dir)
+ os.makedirs(args.json_dir)
+ os.makedirs(args.page_dir)
if args.fixed_id is not None:
filenames = [args.fixed_id]
else:
filenames = get_filenames(args.source_dir)
+ if args.quick:
+ filenames = sorted(set(filenames) - set(quick_used_pages))
print >> sys.stderr, 'Got %d candidates' % len(filenames)
random.shuffle(filenames)
filenames = filenames[:args.limit]
@@ -171,6 +219,14 @@ def main():
print '%2d' % sum(stat.scores),
print 'SCORE %9d (ID %9d) table %2d' % (page_id, meta['id'], table_id),
print stat, stat.scores
+ if args.quick:
+ for table_id, stat in good_table_ids:
+ dump(quick_index, page_id, table_id, stat, args)
+ quick_index += 1
+
+ if args.quick:
+ return
+
random.shuffle(good_tables)
good_tables = good_tables[:args.amount]
if args.dry_run:
@@ -178,22 +234,7 @@ def main():
for i, (page_id, good_table_ids) in enumerate(good_tables):
table_id, stat = random.choice(good_table_ids)
- with open(os.path.join(args.source_dir, '%d.json' % page_id), 'r', 'utf8') as fin:
- meta = json.load(fin)
- meta['tableIndex'] = table_id
- if args.custom_filenames:
- while True:
- try:
- i = int(raw_input('filename for "%s": ' % meta['title']))
- break
- except:
- pass
- with open(os.path.join(json_dir, '%d.json' % i), 'w', 'utf8') as fout:
- json.dump(meta, fout)
- shutil.copy(os.path.join(args.source_dir, '%d.html' % page_id),
- os.path.join(page_dir, '%d.html' % i))
- print '>' * 30, 'Written', i, page_id, table_id, '<' * 30
- print stat
+ dump(i, page_id, table_id, stat, args)
if __name__ == '__main__':
main()
diff --git a/tables/wikipedia-scripts/get-wikipedia-pages.py b/tables/wikipedia-scripts/get-wikipedia-pages.py
index ed1d1f3..99743fc 100755
--- a/tables/wikipedia-scripts/get-wikipedia-pages.py
+++ b/tables/wikipedia-scripts/get-wikipedia-pages.py
@@ -12,7 +12,7 @@ html_parser = HTMLParser.HTMLParser()
def is_sane(content, args):
if content.startswith('#REDIRECT'):
return False
- if args.table and '{|' not in content:
+ if args.table and ('{|' not in content or 'wikitable' not in content):
return False
return True
diff --git a/tables/wikipedia-scripts/table-to-csv.py b/tables/wikipedia-scripts/table-to-csv.py
index afdfdb0..dbe0dbe 100755
--- a/tables/wikipedia-scripts/table-to-csv.py
+++ b/tables/wikipedia-scripts/table-to-csv.py
@@ -1,5 +1,6 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
+"""Convert HTML table into CSV / TSV / pretty-printed table."""
import sys, os, re, argparse, json
from codecs import open
@@ -10,24 +11,40 @@ from itertools import izip_longest
################ Dump CSV
def simple_normalize_text(text):
- return text.replace('\\', '\\\\').replace('"', r'\"').replace('\n', r'\n').replace(u'\xa0', ' ').strip()
+ return text.replace('\\', '\\\\').replace('"', r'\"').replace('\n', r'\\n').replace(u'\xa0', ' ').strip()
def dump_csv(rows, fout):
for row in rows:
- fout.write(','.join('"%s"' % simple_normalize_text(x) for x in row) + '\n')
+ fout.write(','.join('"%s"' % simple_normalize_text(x[1]) for x in row) + '\n')
def tab_normalize_text(text):
- return re.sub(r'\s+', ' ', text.replace('\n', r'\n'), re.U).strip()
+ return re.sub(r'\s+', ' ', text.replace('\\', '\\\\').replace('|', r'\p').replace('\n', r'\n'), re.U).strip()
def dump_tsv(rows, fout):
for row in rows:
- fout.write('\t'.join('%s' % tab_normalize_text(x) for x in row) + '\n')
+ fout.write('\t'.join('%s' % tab_normalize_text(x[1]) for x in row) + '\n')
+
+def table_normalize_text(text):
+ return re.sub(r'\s+', ' ', text, re.U).strip()
+
+def dump_table(rows, fout):
+ widths = defaultdict(int)
+ for row in rows:
+ for i, cell in enumerate(row):
+ widths[i] = max(widths[i], len(table_normalize_text(cell[1])) + 1)
+ for row in rows:
+ fout.write('|')
+ for i, cell in enumerate(row):
+ # wow this is so hacky
+ fout.write((' %-' + str(widths[i]) + 's') % table_normalize_text(cell[1]))
+ fout.write('|')
+ fout.write('\n')
################ More table normalization
def debug_print(stuff):
for x in stuff:
- print >> sys.stderr, [simple_normalize_text(y) for y in x]
+ print >> sys.stderr, [simple_normalize_text(y[1]) for y in x]
def transpose(rows):
cols = []
@@ -38,7 +55,7 @@ def transpose(rows):
try:
col.append(row[i])
except LookupError:
- col.append(None)
+ col.append(('', ''))
cols.append(col)
return cols
@@ -53,7 +70,7 @@ def anti_transpose(cols):
if col[i] is not None:
row.append(col[i])
else:
- row.append('')
+ row.append(('', ''))
rows.append(row)
return rows
@@ -65,7 +82,7 @@ def remove_empty_columns(orig_cols):
"""Remove columns with <= 1 non-empty cells."""
cols = []
for col in orig_cols:
- non_empty = sum((bool(cell) for cell in col), 0)
+ non_empty = sum((bool(cell[1]) for cell in col), 0)
if non_empty >= 2:
cols.append(col)
return cols
@@ -77,9 +94,9 @@ def are_mergeable(col1, col2):
merged = []
for i in xrange(len(col1)):
c1, c2 = col1[i], col2[i]
- if not c1:
+ if not c1[1]:
merged.append(c2)
- elif not c2 or c1 == c2:
+ elif not c2[1] or c1 == c2:
merged.append(c1)
else:
return None
@@ -96,161 +113,93 @@ def merge_similar_columns(orig_cols):
i += 1
return orig_cols
-#### Split columns
+#### Merge header rows
-REGEX_NEWLINE = re.compile(ur'^([^\n]*)\n([^\n]*)$', re.U | re.DOTALL)
-REGEX_NEWLINE_PAREN = re.compile(ur'^(.*)\n(\(.*\))$', re.U | re.DOTALL)
-REGEX_SPLIT_PAREN = re.compile(ur'^(.*) +(\(.*\))$', re.U | re.DOTALL)
-
-def split_multiline_columns(orig_cols):
- """Split columns with newline in each cell."""
- i = 0
- while i < len(orig_cols):
- for regex, threshold in ((REGEX_NEWLINE, 0.5 * len(orig_cols[i])),
- (REGEX_NEWLINE_PAREN, 1),
- (REGEX_SPLIT_PAREN, 2)):
- matches = [regex.match(cell or '') for cell in orig_cols[i]]
- num_matches = sum((bool(match) for match in matches[1:]), 0)
- if num_matches >= threshold:
- splitted = [((None, None) if cell is None
- else (cell, '') if not match
- else (match.group(1), match.group(2)))
- for (cell, match) in zip(orig_cols[i], matches)]
- orig_cols[i:i+1] = [[(None if x[0] is None else x[0].strip()) for x in splitted],
- [(None if x[1] is None else x[1].strip()) for x in splitted]]
- if not orig_cols[i+1][0]:
- orig_cols[i+1][0] = (orig_cols[i][0] or '') + '#n'
- break
- i += 1
- return orig_cols
-
-REGEX_FROM_TO = re.compile(ur'^(.*)[-–—‒―](.*)$', re.U)
-
-def split_from_to_columns(orig_cols):
- """Split columns with pattern '... - ...'."""
- i = 0
- while i < len(orig_cols):
- matches = [REGEX_FROM_TO.match(cell or '') for cell in orig_cols[i]]
- num_matches = sum((bool(match) for match in matches), 0)
- if num_matches > 0.5 * len(orig_cols[i]):
- splitted = [((None, None) if cell is None
- else (cell, '') if not match
- else (match.group(1), match.group(2)))
- for (cell, match) in zip(orig_cols[i], matches)]
- orig_cols[i:i+1] = [[(None if x[0] is None else x[0].strip()) for x in splitted],
- [(None if x[1] is None else x[1].strip()) for x in splitted]]
- if not orig_cols[i+1][0]:
- orig_cols[i+1][0] = orig_cols[i][0]
- orig_cols[i][0] += '#f'
- orig_cols[i+1][0] += '#t'
- i += 1
- return orig_cols
-
-#### Normalize by column
-
-REGEX_NUMBERING = re.compile(ur'^([0-9]+)[.)]$', re.U)
-REGEX_REFERENCE = re.compile(ur'^(.*)[‡^†*]+$', re.U)
-REGEX_PARENS = re.compile(ur'^\((.*)\)$', re.U)
-REGEX_QUOTES = re.compile(ur'^"(.*)"$', re.U)
-
-def normalize_common_punctuations(orig_cols):
- """Normalize some punctuations if they appear a lot in the same row."""
- cols = []
- for col in orig_cols:
- for regex, threshold in ((REGEX_NUMBERING, 0.8 * len(col)),
- (REGEX_REFERENCE, 0.2 * len(col)),
- (REGEX_PARENS, 0.5 * len(col)),
- (REGEX_QUOTES, 0.5 * len(col))):
- matches = [regex.match(cell or '') for cell in col]
- num_matches = sum((not cell or bool(match)
- for (cell, match) in zip(col, matches)), 0)
- if num_matches >= threshold:
- col = [(cell if not match else match.group(1))
- for (cell, match) in zip(col, matches)]
- cols.append(col)
- return cols
-
-#### Normalize by cell
-
-REGEX_COMMA = re.compile(ur'([0-9]),([0-9][0-9][0-9])', re.U)
-
-def remove_commas_from_single_number(x):
- if not x:
- return x
- while REGEX_COMMA.search(x):
- x = REGEX_COMMA.sub(r'\1\2', x)
- return x
-
-def remove_commas_from_numbers(orig_stuff):
- stuff = []
- for slab in orig_stuff:
- stuff.append([remove_commas_from_single_number(x) for x in slab])
- return stuff
+def merge_header_rows(orig_rows):
+ """Merge all header rows together."""
+ header_rows, body_rows = [], []
+ still_header = True
+ for row in orig_rows:
+ if not still_header or any(cell[0] == 'td' for cell in row):
+ still_header = False
+ body_rows.append(row)
+ else:
+ header_rows.append(row)
+ if len(header_rows) < 2 or not body_rows:
+ return orig_rows
+ # Merge header rows with '\n'
+ header_cols = transpose(header_rows)
+ header_row = []
+ for col in header_cols:
+ texts = [None]
+ for cell in col:
+ if cell[1] != texts[-1]:
+ texts.append(cell[1])
+ header_row.append(('th', '\n'.join(texts[1:])))
+ return [header_row] + body_rows
################ Main function
def main():
parser = argparse.ArgumentParser()
- parser.add_argument('-s', '--source-dir', default='wikidump.cache/output',
- help="source directory")
- group = parser.add_mutually_exclusive_group(required=True)
- group.add_argument('-j', '--json',
- help="json metadata file specifying page and table id")
- group.add_argument('-J', '--turk-json',
+ parser.add_argument('-j', '--turk-json',
help="json metadata file from MTurk task")
- group.add_argument('-p', '--page-id', type=int,
- help="page index")
- parser.add_argument('-t', '--table-id', type=int, default=0,
- help="table index (only used for -p / --page-id)")
parser.add_argument('-o', '--outfile',
help="output filename (default = stdout)")
- parser.add_argument('-n', '--normalize', action='count',
- help='degree of normalization')
parser.add_argument('--tsv', action='store_true',
- help='output TSV instead of CSV')
+ help='also print out tsv')
+ parser.add_argument('--human', action='store_true',
+ help='also print out human-readable table')
+ parser.add_argument('--html', action='store_true',
+ help='also print out cleaned html for the table')
+ parser.add_argument('--keep-hidden', action='store_true',
+ help='keep hidden texts as is')
args = parser.parse_args()
+ assert not args.tsv or args.outfile.endswith('.csv')
- if args.json:
- with open(args.json) as fin:
- metadata = json.load(fin)
- args.page_id = metadata['id']
- args.table_id = metadata['tableIndex']
- inhtml = os.path.join(args.source_dir, '%d.html' % args.page_id)
- elif args.turk_json:
- with open(args.turk_json) as fin:
- metadata = json.load(fin)
- args.page_id = metadata['id']
- args.table_id = metadata['tableIndex']
- # The following replacement is pretty hacky:
- inhtml = args.turk_json.replace('-json', '-page').replace('.json', '.html')
- else:
- inhtml = os.path.join(args.source_dir, '%d.html' % args.page_id)
+ with open(args.turk_json) as fin:
+ metadata = json.load(fin)
+
+ # Get the path to the HTML file
+ # This is kind of hacky
+ match = re.match(r'^(?:json|page)/(\d+)-(?:json|page)/(\d+).json$', args.turk_json)
+ batch_id, data_id = match.groups()
+ inhtml = 'page/{}-page/{}.html'.format(batch_id, data_id)
with open(inhtml, 'r', 'utf8') as fin:
- table = Table.get_wikitable(fin.read(), args.table_id, normalization=Table.NORM_DUPLICATE)
+ raw = fin.read()
+ table = Table.get_wikitable(raw, metadata['tableIndex'],
+ normalization=Table.NORM_DUPLICATE,
+ remove_hidden=(not args.keep_hidden))
+ if args.html:
+ raw_table = Table.get_wikitable(raw, metadata['tableIndex'],
+ remove_hidden=False).table
+
rows = table.rows
- if args.normalize >= 1:
- # Remove redundant rows and columns
- rows = remove_full_rowspans(rows)
- cols = transpose(rows)
- cols = remove_empty_columns(cols)
- cols = merge_similar_columns(cols)
- rows = anti_transpose(cols)
- if args.normalize >= 2:
- # Split cells / Normalize texts
- cols = transpose(rows)
- cols = split_multiline_columns(cols)
- cols = normalize_common_punctuations(cols)
- cols = split_from_to_columns(cols)
- cols = remove_commas_from_numbers(cols)
- rows = anti_transpose(cols)
- #debug_print(transpose(rows))
- outputter = dump_tsv if args.tsv else dump_csv
+ # rows = list of columns; column = list of cells; cell = (tag, text)
+ # Remove redundant rows and columns
+ rows = remove_full_rowspans(rows)
+ cols = transpose(rows)
+ cols = remove_empty_columns(cols)
+ cols = merge_similar_columns(cols)
+ rows = anti_transpose(cols)
+ rows = merge_header_rows(rows)
+ # Dump
if not args.outfile:
- outputter(rows, sys.stdout)
+ dump_csv(rows, sys.stdout)
else:
+ stem = re.sub('\.csv$', '', args.outfile)
with open(args.outfile, 'w', 'utf8') as fout:
- outputter(rows, fout)
+ dump_csv(rows, fout)
+ if args.tsv:
+ with open(stem + '.tsv', 'w', 'utf8') as fout:
+ dump_tsv(rows, fout)
+ if args.human:
+ with open(stem + '.table', 'w', 'utf8') as fout:
+ dump_table(rows, fout)
+ if args.html:
+ with open(stem + '.html', 'w', 'utf8') as fout:
+ print >> fout, unicode(raw_table)
if __name__ == '__main__':
main()
diff --git a/tables/wikipedia-scripts/weblib/table.py b/tables/wikipedia-scripts/weblib/table.py
index e39975c..bab1c9f 100644
--- a/tables/wikipedia-scripts/weblib/table.py
+++ b/tables/wikipedia-scripts/weblib/table.py
@@ -7,7 +7,9 @@ Get statistics about a table and convert it to CSV.
import sys, os, re, json
from codecs import open
from collections import defaultdict
-from bs4 import BeautifulSoup
+from bs4 import BeautifulSoup as BeautifulSoupOriginal
+def BeautifulSoup(markup=""):
+ return BeautifulSoupOriginal(markup, 'html.parser')
class Table(object):
NORM_NONE = 0
@@ -21,14 +23,11 @@ class Table(object):
self.table = table
if remove_hidden:
self.remove_hidden()
- if normalization == Table.NORM_NONE:
- self.get_cells()
- elif normalization == Table.NORM_CORNER:
+ if normalization == Table.NORM_CORNER:
self.normalize_table()
- self.get_cells()
elif normalization == Table.NORM_DUPLICATE:
self.normalize_table(deep=True)
- self.get_cells()
+ self.get_cells()
@staticmethod
def get_wikitable(raw_html, index=None, **kwargs):
@@ -53,11 +52,12 @@ class Table(object):
tag.extract()
def get_cells(self):
+ """Each cell is (tag, text)"""
self.rows, self.cells = [], []
for x in self.table.find_all('tr', recursive=False):
row = []
for y in x.find_all(['th', 'td'], recursive=False):
- row.append(y.text.strip())
+ row.append((y.name, y.text.strip()))
self.rows.append(row)
self.cells.extend(row)
self.num_rows = len(self.rows)
@@ -82,7 +82,7 @@ class Table(object):
def get_cloned_cell(self, cell, rowspan=1, deep=False):
if deep:
# Hacky but works
- return BeautifulSoup(unicode(cell)).body.contents[0]
+ return BeautifulSoup(unicode(cell)).contents[0]
tag = Table.SOUP.new_tag(cell.name)
if rowspan > 1:
tag['rowspan'] = rowspan
@@ -152,4 +152,4 @@ def test():
print table.table
if __name__ == '__main__':
- test()
+ test_wiki()
diff --git a/tables/x b/tables/x
index bfa9d06..c561759 100755
--- a/tables/x
+++ b/tables/x
@@ -16,6 +16,7 @@ if [ "$operation" == "up" ]; then
ln -s ../module-classes.txt xxx/
ln -s ../tables/grammars xxx/
ln -s ../run xxx/
+ { hostname; readlink -f .; date; git log -1; git status; } > xxx/git-hash
cl upload -L -n stuff xxx
rm -rf xxx
exit 0