From 9bd5b7bb2aa97c4157a61e772fcb15157792d400 Mon Sep 17 00:00:00 2001 From: gauravbaruah Date: Sun, 7 May 2017 22:28:35 -0400 Subject: [PATCH] Corpus idf (#23) as part of sourcing-IDF-from-index and e2e experiments. --- .gitignore | 3 +- README.md | 2 +- idf_baseline/README.md | 7 +- idf_baseline/experimental_settings.py | 157 +++++++++++++++++ idf_baseline/pom.xml | 8 +- idf_baseline/qa-data-only-idf.py | 95 +++++++++-- .../main/java/ai/castor/idf/FetchTermIDF.java | 120 +++++++++++++ sm_cnn/README.md | 6 +- sm_cnn/bridge.py | 43 ++++- sm_cnn/experimental_settings.py | 160 ++++++++++++++++++ sm_cnn/external_features.py | 63 ++++++- sm_cnn/main.py | 19 ++- sm_cnn/make_run.py | 51 ------ sm_cnn/train.py | 11 +- sm_cnn/utils.py | 74 ++++++-- 15 files changed, 707 insertions(+), 112 deletions(-) create mode 100644 idf_baseline/experimental_settings.py create mode 100644 idf_baseline/src/main/java/ai/castor/idf/FetchTermIDF.java create mode 100644 sm_cnn/experimental_settings.py delete mode 100644 sm_cnn/make_run.py diff --git a/.gitignore b/.gitignore index 00741cb..caed9fa 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ .DS_Store -.idea/ \ No newline at end of file +.idea/ +*idfsim diff --git a/README.md b/README.md index 9221d19..d4e3975 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ Pytorch deep learning models. -1. [SM model](./sm_model/): Similarity between question and candidate answers. +1. [SM model](./sm_cnn/): Similarity between question and candidate answers. ## Setting up Pytorch diff --git a/idf_baseline/README.md b/idf_baseline/README.md index d8d72a4..09fe9f9 100644 --- a/idf_baseline/README.md +++ b/idf_baseline/README.md @@ -92,7 +92,7 @@ mvn clean package appassembler:assemble Run the following command to score each answer with an IDF value: ``` -sh target/appassembler/bin/GetIDF -index ~/large-local-work/indices/index.wikipedia.pos.docvectors -config ../../data/WikiQA/test -output WikiQA.test.idfsim +sh target/appassembler/bin/GetIDFSumSimilarity -index ~/large-local-work/indices/index.wikipedia.pos.docvectors -config ../../data/WikiQA/test -output WikiQA.test.idfsim ``` The above command will create a run file in the `trec_eval` format and a qrel file at a location specified by `-output`. @@ -154,5 +154,10 @@ python qa-data-only-idf.py ../../data/WikiQA WikiQA ``` Evaluate these using step 2. +The same script can now also be used to comput idf sum similarity based on corpus idf statistics +``` +python qa-data-only-idf.py ../../data/TrecQA TrecQA --index-for-corpusIDF ../../data/indices/index.qadata.pos.docvectors.keepstopwords/ +``` + ### Baseline results Baseline results are saved in ``Castor/baseline_results.tsv`` \ No newline at end of file diff --git a/idf_baseline/experimental_settings.py b/idf_baseline/experimental_settings.py new file mode 100644 index 0000000..3b88acb --- /dev/null +++ b/idf_baseline/experimental_settings.py @@ -0,0 +1,157 @@ +# The following 5 conditions vary +# idf_source, stopwords_and_stemming, punctuation, words-with-hyphens +import argparse +import itertools +import shlex +import subprocess + +class Setting(object): + def __init__(self, label, value_flag_map): + self.label = label + self.choice_flags = value_flag_map + + def get_settings(self): + return self.choice_flags.keys() + + def get_choice(self, setting): + return self.choice_flags[setting] + + def get_options(self): + options = [] + for key in self.choice_flags.keys(): + options.append("{}:{}".format(self.label, key)) + return options + +class Experiments(object): + def __init__(self, qa_dataset): + self.settings = {} + self.combinations = [] + self.qa_data = qa_dataset + self.cmd_root = "python qa-data-only-idf.py {} run".format(self.qa_data) + self.eval_cmd_root = "../../Anserini/eval/trec_eval.9.0/trec_eval -m map -m recip_rank -m bpref" + self.rbp_cmd_root = "rbp_eval" + + def add_setting(self, setting): + self.settings[setting.label] = setting + self._setup_combinations() + + + def _setup_combinations(self): + all_settings = [] + for setting in self.settings.values(): + all_settings.append(setting.get_options()) + + self.combinations = [] + for c in itertools.product(*all_settings): + self.combinations.append(c) + + def _run_cmd(self, cmd): + pargs = shlex.split(cmd) + p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE, \ + bufsize=1, universal_newlines=True) + pout, perr = p.communicate() + return pout, perr + + def _run_eval(self): + for split in ['train-all', 'raw-dev', 'raw-test']: + cmd = '{} {}/{}.qrel run.{}.idfsim'.format(self.eval_cmd_root, + self.qa_data, split, split) + out, err = self._run_cmd(cmd) + print(split, '------' ) + # trec_eval scores + metrics = [] + scores = [] + for line in str(out).split('\n'): + if not line.strip().split(): + continue + fields = line.strip().split() + metrics.append(fields[0]) + scores.append(fields[-1]) + + # rbp_eval scores + cmd = '{} {}/{}.qrel run.{}.idfsim'.format(self.rbp_cmd_root, + self.qa_data, split, split) + out, err = self._run_cmd(cmd) + for line in str(out).split('\n'): + if not line.startswith('p= 0.50'): + continue + metrics.append('rbp_p0.5') + scores.append(' '.join(line.strip().split()[-2:])) + + print('\t'.join(metrics)) + print('\t'.join(scores)) + + + def run(self, indices): + """ + runs a particular combination of settings + """ + for ci in indices: + combo = self.combinations[ci] + print(combo) + cmd_args = [] + for setting_choice in combo: + setting, choice = setting_choice.split(':') + cmd_args.append(self.settings[setting].choice_flags[choice]) + cmd = '{} {}'.format(self.cmd_root, ' '.join(cmd_args)) + print(cmd) + out, err = self._run_cmd(cmd) + self._run_eval() + + def run_all(self): + """ + runs all experiments + """ + pass + + def list_settings(self): + """ + lists all settings + """ + for c in enumerate(self.combinations): + print(c) + print("--run X Y Z to run combinations number X Y Z") + + +if __name__ == "__main__": + ap = argparse.ArgumentParser(description="Lists exerimental settings and runs experiments") + ap.add_argument("--list", help="lists all available experimental settings combinations", + action="store_true") + ap.add_argument("--run", help='runs experimenal setting combination NUMBER(s)', + nargs="+", type=int) + ap.add_argument("--runall", help="runs all experiments in order", action="store_true") + ap.add_argument("index_path", help="required for some combination of experiments") + ap.add_argument('qa_data', help="path to the QA dataset", + choices=['../../data/TrecQA', '../../data/WikiQA']) + + args = ap.parse_args() + + experiments = Experiments(args.qa_data) + + experiments.add_setting(Setting('idf_source', { + 'qa-data':'', + 'corpus-index': '--index-for-corpusIDF {}'.format(args.index_path) + })) + + experiments.add_setting(Setting('stop_stem', { + 'yes':'--stop-and-stem', + 'no': '' + })) + + experiments.add_setting(Setting('punctuation', { + 'keep': '', + 'remove': '--stop-punct' + })) + + experiments.add_setting(Setting('dash_words', { + 'keep': '', + 'split': '--dash-split' + })) + + experiments.list_settings() + + if args.run: + experiments.run(args.run) + + if args.runall: + experiments.run_all() diff --git a/idf_baseline/pom.xml b/idf_baseline/pom.xml index e8a6e1c..9ed9bdd 100644 --- a/idf_baseline/pom.xml +++ b/idf_baseline/pom.xml @@ -46,12 +46,15 @@ ai.castor.idf.IDFScorer - GetIDF + GetIDFSumSimilarity + + + ai.castor.idf.FetchTermIDF + FetchTermIDF - maven-assembly-plugin @@ -65,7 +68,6 @@ - diff --git a/idf_baseline/qa-data-only-idf.py b/idf_baseline/qa-data-only-idf.py index 4b72ada..b49a7ee 100644 --- a/idf_baseline/qa-data-only-idf.py +++ b/idf_baseline/qa-data-only-idf.py @@ -1,23 +1,48 @@ import argparse import os +import sys +import re import numpy as np from collections import defaultdict import string +import subprocess +import shlex import nltk -nltk.download('stopwords') +nltk.download('stopwords', quiet=True) from nltk.stem.porter import PorterStemmer from nltk.corpus import stopwords -def read_in_data(datapath, set_name, file, stop_and_stem=False): +def read_in_data(datapath, set_name, file, stop_and_stem=False, stop_punct=False, dash_split=False): data = [] with open(os.path.join(datapath, set_name, file)) as inf: data = [line.strip() for line in inf.readlines()] + + if dash_split: + def split_hyphenated_words(sentence): + rtokens = [] + for term in sentence.split(): + for t in term.split('-'): + if t: + rtokens.append(t) + return ' '.join(rtokens) + data = [split_hyphenated_words(sentence) for sentence in data] + + if stop_punct: + regex = re.compile('[{}]'.format(re.escape(string.punctuation))) + def remove_punctuation(sentence): + rtokens = [] + for term in sentence.split(): + for t in regex.sub(' ', term).strip().split(): + if t: + rtokens.append(t) + return ' '.join(rtokens) + data = [remove_punctuation(sentence) for sentence in data] + if stop_and_stem: stemmer = PorterStemmer() stoplist = set(stopwords.words('english')) - stoplist.update(set(string.punctuation)) def stop_stem(sentence): return ' '.join([stemmer.stem(word) for word in sentence.split() \ if word not in stoplist]) @@ -25,16 +50,47 @@ def read_in_data(datapath, set_name, file, stop_and_stem=False): return data -def compute_idfs(data): +def compute_idfs(data, dash_split=False): term_idfs = defaultdict(float) for doc in list(data): for term in list(set(doc.split())): + if dash_split: + assert('-' not in term) term_idfs[term] += 1.0 N = len(data) for term, n_t in term_idfs.items(): term_idfs[term] = np.log(N/(1+n_t)) return term_idfs +def fetch_idfs_from_index(data, dash_split, indexPath): + regex = re.compile('[{}]'.format(re.escape(string.punctuation))) + term_idfs = defaultdict(float) + all_terms = set([term for doc in list(data) for term in doc.split()]) + with open('dataset.vocab', 'w') as vf: + for term in list(all_terms): + if dash_split: + assert('-' not in term) + print(term, file=vf) + + fetchIDF_cmd = \ + "sh ../idf_baseline/target/appassembler/bin/FetchTermIDF -index {} -vocabFile {}".\ + format(indexPath, 'dataset.vocab') + pargs = shlex.split(fetchIDF_cmd) + p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE, \ + bufsize=1, universal_newlines=True) + pout, perr = p.communicate() + + lines = str(pout).split('\n') + for line in lines: + if not line: + continue + fields = line.strip().split("\t") + term, weight = fields[0], fields[-1] + term_idfs[term] = float(weight) + + for line in str(perr).split('\n'): + print('Warning: '+line) + return term_idfs def compute_idf_sum_similarity(questions, answers, term_idfs): # compute IDF sums for common_terms @@ -51,7 +107,7 @@ def compute_idf_sum_similarity(questions, answers, term_idfs): def write_out_idf_sum_similarities(qids, questions, answers, term_idfs, outfile, dataset): - with open(outfile, 'w') as outf: + with open(outfile, 'w') as outf: idf_sum_similarity = compute_idf_sum_similarity(questions, answers, term_idfs) old_qid = 0 docid_c = 0 @@ -59,7 +115,7 @@ def write_out_idf_sum_similarities(qids, questions, answers, term_idfs, outfile, if qids[i] != old_qid and dataset.endswith('WikiQA'): docid_c = 0 old_qid = qids[i] - print('{} 0 {} 0 {} data_only_idfbaseline'.format(qids[i], docid_c, + print('{} 0 {} 0 {} idfbaseline'.format(qids[i], docid_c, idf_sum_similarity[i]), file=outf) docid_c += 1 @@ -74,6 +130,10 @@ if __name__ == "__main__": ap.add_argument('--ignore-test', help="does not consider test data when computing IDF of terms", action="store_true") ap.add_argument("--stop-and-stem", help='performs stopping and stemming', action="store_true") + ap.add_argument("--stop-punct", help='removes punctuation', action="store_true") + ap.add_argument("--dash-split", help="split words containing hyphens", action="store_true") + ap.add_argument("--index-for-corpusIDF", help="fetches idf from Index. provide index path. will\ + generate a vocabFile") args = ap.parse_args() @@ -82,14 +142,20 @@ if __name__ == "__main__": if args.qa_data.endswith('TrecQA'): train_data, dev_data, test_data = 'train-all', 'raw-dev', 'raw-test' - train_que = read_in_data(args.qa_data, train_data, 'a.toks', args.stop_and_stem) - train_ans = read_in_data(args.qa_data, train_data, 'b.toks', args.stop_and_stem) + train_que = read_in_data(args.qa_data, train_data, 'a.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) + train_ans = read_in_data(args.qa_data, train_data, 'b.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) - dev_que = read_in_data(args.qa_data, dev_data, 'a.toks', args.stop_and_stem) - dev_ans = read_in_data(args.qa_data, dev_data, 'b.toks', args.stop_and_stem) + dev_que = read_in_data(args.qa_data, dev_data, 'a.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) + dev_ans = read_in_data(args.qa_data, dev_data, 'b.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) - test_que = read_in_data(args.qa_data, test_data, 'a.toks', args.stop_and_stem) - test_ans = read_in_data(args.qa_data, test_data, 'b.toks', args.stop_and_stem) + test_que = read_in_data(args.qa_data, test_data, 'a.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) + test_ans = read_in_data(args.qa_data, test_data, 'b.toks', + args.stop_and_stem, args.stop_punct, args.dash_split) all_data = train_que + dev_que + train_ans + dev_ans @@ -98,7 +164,10 @@ if __name__ == "__main__": all_data += test_que # compute inverse document frequencies for terms - term_idfs = compute_idfs(set(all_data)) + if not args.index_for_corpusIDF: + term_idfs = compute_idfs(set(all_data), args.dash_split) + else: + term_idfs = fetch_idfs_from_index(set(all_data), args.dash_split, args.index_for_corpusIDF) # write out in trec_eval format write_out_idf_sum_similarities(read_in_data(args.qa_data, train_data, 'id.txt'), diff --git a/idf_baseline/src/main/java/ai/castor/idf/FetchTermIDF.java b/idf_baseline/src/main/java/ai/castor/idf/FetchTermIDF.java new file mode 100644 index 0000000..4f48d81 --- /dev/null +++ b/idf_baseline/src/main/java/ai/castor/idf/FetchTermIDF.java @@ -0,0 +1,120 @@ +/** + * Anserini: An information retrieval toolkit built on Lucene + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package ai.castor.idf; + +import org.apache.lucene.analysis.Analyzer; +import org.apache.lucene.analysis.CharArraySet; +import org.apache.lucene.analysis.StopFilter; +import org.apache.lucene.analysis.core.WhitespaceAnalyzer; +import org.apache.lucene.analysis.en.EnglishAnalyzer; +import org.apache.lucene.index.DirectoryReader; +import org.apache.lucene.index.IndexReader; +import org.apache.lucene.index.Term; +import org.apache.lucene.queryparser.classic.ParseException; +import org.apache.lucene.queryparser.classic.QueryParser; +import org.apache.lucene.search.*; +import org.apache.lucene.search.similarities.ClassicSimilarity; +import org.apache.lucene.store.FSDirectory; +import org.kohsuke.args4j.*; + +import java.io.*; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.*; + +public class FetchTermIDF { + + public static class Args { + @Option(name = "-index", metaVar="[path]", required = true, usage = "path to Lucene index") + public String index; + + @Option(name = "-terms", usage = "space separated list of terms to get idf for.") + public String terms; + + @Option(name = "-vocabFile", usage = "file with one term per line") + public String vocabFile; + } + + private final IndexReader reader; + private final FSDirectory directory; + + public static final String FIELD_BODY = "contents"; + + public FetchTermIDF(FetchTermIDF.Args args) throws Exception { + Path indexPath = Paths.get(args.index); + if (!Files.exists(indexPath) || !Files.isDirectory(indexPath) || !Files.isReadable(indexPath)) { + throw new IllegalArgumentException(args.index + " does not exist, is not a directory, or is not readable"); + } + this.directory = FSDirectory.open(indexPath); + this.reader = DirectoryReader.open(this.directory); // #changed + } + + public double getTermIDF(String term) throws ParseException { + Analyzer analyzer = new EnglishAnalyzer(CharArraySet.EMPTY_SET); + QueryParser qp = new QueryParser(FIELD_BODY, analyzer); + ClassicSimilarity similarity = new ClassicSimilarity(); + + String esTerm = qp.escape(term); + double termIDF = 0.0; + try { + TermQuery q = (TermQuery) qp.parse(esTerm); + Term t = q.getTerm(); + termIDF = similarity.idf(reader.docFreq(t), reader.numDocs()); + + System.out.println(term + '\t' + esTerm + '\t' + q + '\t' + t + '\t' + termIDF); + } catch (Exception e) { + System.err.println("Exception in fetching IDF(" + term + "): " + e.toString()); + } + return termIDF; + } + + public static void main(String[] args) throws Exception { + Args qaArgs = new Args(); + CmdLineParser parser = new CmdLineParser(qaArgs, ParserProperties.defaults().withUsageWidth(90)); + + try { + parser.parseArgument(args); + } catch (CmdLineException e) { + System.err.println(e.getMessage()); + parser.printUsage(System.err); + System.err.println("Example: FetchTermIDF" + parser.printExample(OptionHandlerFilter.REQUIRED)); + return; + } + + if (qaArgs.terms == null && qaArgs.vocabFile == null) { + System.out.println("Required one of -vocabFile or -terms arguments"); + return; + } + + FetchTermIDF idfFetcher = new FetchTermIDF(qaArgs); + List termsList = null; + if (qaArgs.terms != null && !qaArgs.terms.isEmpty()) { + termsList = Arrays.asList(qaArgs.terms.split("\\s+")); + for(String t: termsList) { + idfFetcher.getTermIDF(t); + } + } else { + try(BufferedReader br = new BufferedReader(new FileReader(qaArgs.vocabFile))) { + for(String line; (line = br.readLine()) != null; ) { + idfFetcher.getTermIDF(line); + } + // line is not visible here. + } + } + } + +} \ No newline at end of file diff --git a/sm_cnn/README.md b/sm_cnn/README.md index b63d99d..99c7e38 100644 --- a/sm_cnn/README.md +++ b/sm_cnn/README.md @@ -28,9 +28,11 @@ git clone https://github.com/castorini/Castor.git This should generate: ``` ├── Castor +│   ├── README.md │   ├── idf_baseline │   ├── kim_cnn -│   └── sm_cnn +│   ├── simple_qa_rnn +│   └── sm_cnn/ ├── data │   ├── README.md │   ├── TrecQA/ @@ -65,7 +67,7 @@ cd .. To train the S&M model on TrecQA ``` -python main.py ../../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor +python main.py ../../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor --paper-ext-features ``` The final model will be saved to ```../../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor``` diff --git a/sm_cnn/bridge.py b/sm_cnn/bridge.py index 07d88d8..41761d0 100644 --- a/sm_cnn/bridge.py +++ b/sm_cnn/bridge.py @@ -4,6 +4,8 @@ import sys from collections import Counter import argparse +import re +import string import numpy as np import torch from nltk.tokenize import TreebankWordTokenizer @@ -25,6 +27,8 @@ class SMModelBridge(object): # load model self.model = model.QAModel.load(model_file) + self.model_file = model_file + # load vectors self.vec_dim = self._preload_cached_embeddings(word_embeddings_cache_file) self.unk_term_vec = np.random.uniform(-0.25, 0.25, self.vec_dim) @@ -48,8 +52,34 @@ class SMModelBridge(object): def parse(self, sentence): s_toks = TreebankWordTokenizer().tokenize(sentence) - s_str = ' '.join(s_toks).lower() - return s_str + sentence = ' '.join(s_toks).lower() + + model_input_args = self.model_file.split('.') + punctuation = model_input_args[-3].split('-')[1] + dash_words = model_input_args[-2].split('_')[1] + + if dash_words == "split": + def split_hyphenated_words(sentence): + rtokens = [] + for term in sentence.split(): + for t in term.split('-'): + if t: + rtokens.append(t) + return ' '.join(rtokens) + sentence = split_hyphenated_words(sentence) + + if punctuation == "remove": + regex = re.compile('[{}]'.format(re.escape(string.punctuation))) + def remove_punctuation(sentence): + rtokens = [] + for term in sentence.split(): + for t in regex.sub(' ', term).strip().split(): + if t: + rtokens.append(t) + return ' '.join(rtokens) + sentence = remove_punctuation(sentence) + + return sentence def make_input_matrix(self, sentence): @@ -86,9 +116,16 @@ class SMModelBridge(object): term_idfs = json.loads(idf_json) term_idfs = dict((k, float(v)) for k, v in term_idfs.items()) + for term in question.split(): + if term not in term_idfs: + term_idfs[term] = 0.0 + for answer in answers: answer = self.parse(answer) - + for term in answer.split(): + if term not in term_idfs: + term_idfs[term] = 0.0 + overlap = compute_overlap([question], [answer]) idf_weighted_overlap = compute_idf_weighted_overlap([question], [answer], term_idfs) overlap_no_stopwords =\ diff --git a/sm_cnn/experimental_settings.py b/sm_cnn/experimental_settings.py new file mode 100644 index 0000000..77a4b8a --- /dev/null +++ b/sm_cnn/experimental_settings.py @@ -0,0 +1,160 @@ +# The following 5 conditions vary +# idf_source, stopwords_and_stemming, punctuation, words-with-hyphens +import argparse +import itertools +import shlex +import subprocess + +class Setting(object): + def __init__(self, label, value_flag_map): + self.label = label + self.choice_flags = value_flag_map + + def get_settings(self): + return self.choice_flags.keys() + + def get_choice(self, setting): + return self.choice_flags[setting] + + def get_options(self): + options = [] + for key in self.choice_flags.keys(): + options.append("{}:{}".format(self.label, key)) + return options + +class Experiments(object): + def __init__(self, qa_dataset, word_embeddings_file): + self.settings = {} + self.combinations = [] + self.qa_data = qa_dataset + self.w2v_file = word_embeddings_file + #self.cmd_root = "python main.py --dataset_folder {} --word_vectors_file {} --run-name-prefix run --epochs 1 --num_conv_filters 5".format(self.qa_data, self.w2v_file) + self.cmd_root = "python main.py --dataset_folder {} --word_vectors_file {} --run-name-prefix run --paper-ext-feats".format(self.qa_data, self.w2v_file) + self.eval_cmd_root = "../../Anserini/eval/trec_eval.9.0/trec_eval -m map -m recip_rank -m bpref" + self.rbp_cmd_root = "rbp_eval" + + def add_setting(self, setting): + self.settings[setting.label] = setting + self._setup_combinations() + + + def _setup_combinations(self): + all_settings = [] + for setting in self.settings.values(): + all_settings.append(setting.get_options()) + + self.combinations = [] + for c in itertools.product(*all_settings): + self.combinations.append(c) + + def _run_cmd(self, cmd): + pargs = shlex.split(cmd) + p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE, \ + bufsize=1, universal_newlines=True) + pout, perr = p.communicate() + return pout, perr + + def _run_eval(self): + for split in ['raw-dev', 'raw-test']: + cmd = '{} {}/{}.qrel run.{}.smrun'.format(self.eval_cmd_root, + self.qa_data, split, split) + out, err = self._run_cmd(cmd) + print(split, '------' ) + # trec_eval scores + metrics = [] + scores = [] + for line in str(out).split('\n'): + if not line.strip().split(): continue + fields = line.strip().split() + metrics.append(fields[0]) + scores.append(fields[-1]) + + # rbp_eval scores + cmd = '{} {}/{}.qrel run.{}.smrun'.format(self.rbp_cmd_root, + self.qa_data, split, split) + out, err = self._run_cmd(cmd) + for line in str(out).split('\n'): + if not line.startswith('p= 0.50'): continue + metrics.append('rbp_p0.5') + scores.append(' '.join(line.strip().split()[-2:])) + + print('\t'.join(metrics)) + print('\t'.join(scores)) + + + def run(self, indices): + """ + runs a particular combination of settings + """ + for ci in indices: + combo = self.combinations[ci] + print(combo) + cmd_args = [] + # set model name + model_name = 'sm_cnn.' + for setting_choice in combo: + setting, choice = setting_choice.split(':') + model_name += '{}-{}.'.format(setting, choice) + cmd_args.append(self.settings[setting].choice_flags[choice]) + model_name += 'model' + cmd = '{} {} {}'.format(self.cmd_root, ' '.join(cmd_args), model_name) + print(cmd) + out, err = self._run_cmd(cmd) + with open(model_name + '.log', 'w') as lf: + print('---------- OUT ------------', file=lf) + print(out, file=lf) + print('---------- ERR ------------', file=lf) + print(err, file=lf) + self._run_eval() + + def run_all(self): + """ + runs all experiments + """ + pass + + def list_settings(self): + """ + lists all settings + """ + for c in enumerate(self.combinations): + print(c) + print("--run X Y Z to run combinations number X Y Z") + + +if __name__ == "__main__": + ap = argparse.ArgumentParser(description="Lists exerimental settings and runs experiments") + ap.add_argument("--list", help="lists all available experimental settings combinations", + action="store_true") + ap.add_argument("--run", help='runs experimenal setting combination NUMBER(s)', + nargs="+", type=int) + ap.add_argument("indexPath", help="required for some combination of experiments") + ap.add_argument('qa_data', help="path to the QA dataset", + choices=['../../data/TrecQA', '../../data/WikiQA']) + ap.add_argument('word_embeddings_file', help="the word embeddings file") + + args = ap.parse_args() + + experiments = Experiments(args.qa_data, args.word_embeddings_file) + + experiments.add_setting(Setting('idf_source', { + 'qa-data':'', + 'corpus-index': '--index-for-corpusIDF {}'.format(args.indexPath) + })) + + experiments.add_setting(Setting('punctuation', { + 'keep': '', + 'remove': '--stop-punct' + })) + + experiments.add_setting(Setting('dash_words', { + 'keep': '', + 'split': '--dash-split' + })) + + experiments.list_settings() + + if args.run: + experiments.run(args.run) + + diff --git a/sm_cnn/external_features.py b/sm_cnn/external_features.py index dcc5b96..83e5191 100644 --- a/sm_cnn/external_features.py +++ b/sm_cnn/external_features.py @@ -1,16 +1,19 @@ -# module to compute various external features for the sm_model. +# module to compute various external features for the sm cnn model. # TODO: add more external features like: # word mover distance, cosine sim in tf.idf space, cosine sim in word embedding space # overlap based on parts of speech: noun, verb, adj (POS tag) # word embedding cosine sim based on part of speech: noun, verb, adj - +import sys +import os +import shlex +import subprocess import string from collections import defaultdict import numpy as np import nltk -nltk.download('stopwords') +nltk.download('stopwords', quiet=True) from nltk.stem.porter import PorterStemmer from nltk.corpus import stopwords @@ -49,11 +52,50 @@ def get_qadata_only_idf(all_data): term_idfs[term] = np.log(N/(1+n_t)) return term_idfs -def get_source_corpus_idf(all_data): +def get_source_corpus_idf(all_data, path_to_index): """ fetches idf weights from source corpus (disks1-5+aquaint|wikipedia) index, for all the qa pairs """ - pass + # first run maven to build ../idf_baseline/FetchTermIDF + maven_cmd = "mvn -f ../idf_baseline/pom.xml clean package appassembler:assemble" + pargs = shlex.split(maven_cmd) + p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE, \ + bufsize=1, universal_newlines=True) + pout, perr = p.communicate() + # if build failure, exit with message + if "BUILD FAILURE" in pout or "BUILD FAILURE" in perr: + print("\nERROR: Could not build ../idf_baseline/FetchTermIDF. Fix build errors before proceeding") + print("$ cd ../idf_baseline") + print("$ mvn clean package appassembler:assemble") + sys.exit(0) + + if not type(all_data) is list: + all_data = list(all_data) + term_idfs = defaultdict(float) + all_terms = set([term for doc in all_data for term in doc.split()]) + with open('dataset.vocab', 'w') as vf: + for term in list(all_terms): + print(term, file=vf) + + fetchIDF_cmd = \ + "sh ../idf_baseline/target/appassembler/bin/FetchTermIDF -index {} -vocabFile {}".\ + format(path_to_index, 'dataset.vocab') + pargs = shlex.split(fetchIDF_cmd) + p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE, \ + bufsize=1, universal_newlines=True) + pout, perr = p.communicate() + + lines = str(pout).split('\n') + for line in lines: + if not line: + continue + fields = line.strip().split("\t") + term, weight = fields[0], fields[-1] + term_idfs[term] = float(weight) + + for line in str(perr).split('\n'): + print('Warning: '+line) + return term_idfs def compute_overlap(questions, answers): """ @@ -83,7 +125,7 @@ def compute_idf_weighted_overlap(questions, answers, idf_weights): return np.array(overlap_scores) -def set_external_features_as_per_paper(trainer): +def set_external_features_as_per_paper(trainer, corpus_index=None): """ computes external features as per the paper AND saves them into trainer """ @@ -95,7 +137,9 @@ def set_external_features_as_per_paper(trainer): all_answers.extend(answers) all_data = set(all_questions + all_answers) - idf_weights = get_qadata_only_idf(list(all_data)) + print('corpus_index', corpus_index) + idf_weights = get_qadata_only_idf(list(all_data)) if not corpus_index else \ + get_source_corpus_idf(list(all_data), corpus_index) external_features = {} @@ -122,7 +166,7 @@ def set_external_features_as_per_paper(trainer): return external_features -def set_external_features_as_per_paper_and_stem(trainer): +def set_external_features_as_per_paper_and_stem(trainer, corpus_index=None): """ computes external features as per the paper but performs stemming before computing IDF. features are saved into the trainer.data_splits @@ -143,7 +187,8 @@ def set_external_features_as_per_paper_and_stem(trainer): return ' '.join([stemmer.stem(word) if word not in stoplist else word \ for word in sentence.split()]) all_but_stopwords_stemmed = [stem_non_stop_words(sentence) for sentence in list(all_data)] - idf_weights = get_qadata_only_idf(all_but_stopwords_stemmed) + idf_weights = get_qadata_only_idf(all_but_stopwords_stemmed) if not corpus_index else \ + get_source_corpus_idf(all_but_stopwords_stemmed, corpus_index) external_features = {} diff --git a/sm_cnn/main.py b/sm_cnn/main.py index da2832d..39fa302 100644 --- a/sm_cnn/main.py +++ b/sm_cnn/main.py @@ -78,7 +78,7 @@ if __name__ == "__main__": formatter_class=argparse.ArgumentDefaultsHelpFormatter) ap.add_argument('model_outfile', help='file to save final model') ap.add_argument('--word_vectors_file', \ - help='NOTE: a cache will be created for faster loading for word vectors', + help='NOTE: a cache will be created for faster loading for word vectors',\ default="../../data/word2vec/aquaint+wiki.txt.gz.ndim=50.bin") ap.add_argument('--dataset_folder', help='directory containing train, dev, test sets', \ default="../../data/TrecQA") @@ -86,7 +86,7 @@ if __name__ == "__main__": # external features related arguments ap.add_argument('--no-ext-feats', action="store_true", \ help="will not include external features in the model") - ap.add_argument('--paper-ext-feats', action="store_true", \ + ap.add_argument('--paper-ext-feats', action="store_true", default=True, \ help="external features as per the paper") ap.add_argument('--paper-ext-feats-stem', action="store_true", \ help="external features as per the paper") @@ -101,7 +101,7 @@ if __name__ == "__main__": ap.add_argument('--mom', help='SGD Momentum', default=0.0, type=float) ap.add_argument('--train', help='switches to train set', action="store_true") # epoch related arguments - ap.add_argument('--epochs', type=int, default=25, help="number of trainin epochs") + ap.add_argument('--epochs', type=int, default=25, help="number of training epochs") ap.add_argument('--patience', type=int, default=5, \ help="if there is no appreciable change in model after epochs, then stop") # debugging arguments @@ -114,6 +114,10 @@ if __name__ == "__main__": help='runs test on each epoch to track final performance') ap.add_argument("--skip-training", help="will load pre-trained model", action="store_true") ap.add_argument("--run-name-prefix", help="will output train|dev|test runs with provided prefix") + ap.add_argument("--stop-punct", help='removes punctuation', action="store_true") + ap.add_argument("--dash-split", help="split words containing hyphens", action="store_true") + ap.add_argument("--index-for-corpusIDF", help="fetches idf from Index. provide index path. will\ + generate a vocabFile") args = ap.parse_args() @@ -144,12 +148,15 @@ if __name__ == "__main__": # TODO: remember to update args.* in testing loop below if args.paper_ext_feats: logger.info("--paper-ext-feats") - ext_feats_for_splits = set_external_features_as_per_paper(trainer) + ext_feats_for_splits = \ + set_external_features_as_per_paper(trainer, args.index_for_corpusIDF) # ^^ we are saving the features to be used while testing at the end of training elif args.paper_ext_feats_stem: logger.info("--paper-ext-feats-stem") - ext_feats_for_splits = set_external_features_as_per_paper_and_stem(trainer) + ext_feats_for_splits = \ + set_external_features_as_per_paper_and_stem(trainer, args.index_for_corpusIDF) + if not args.skip_training: best_map = 0.0 best_model = 0 @@ -187,7 +194,7 @@ if __name__ == "__main__": trained_model = QAModel.load(args.model_outfile) evaluator = Trainer(trained_model, args.eta, args.mom, args.no_loss_reg, vec_dim) - for split in [test_set, dev_set, train_set]: + for split in [test_set, dev_set]: evaluator.load_input_data(args.dataset_folder, cache_file, None, None, split) if args.paper_ext_feats or args.paper_ext_feats_stem: evaluator.data_splits[split][-1] = ext_feats_for_splits[split] diff --git a/sm_cnn/make_run.py b/sm_cnn/make_run.py deleted file mode 100644 index b1044d6..0000000 --- a/sm_cnn/make_run.py +++ /dev/null @@ -1,51 +0,0 @@ -import argparse -import os -from model import QAModel -from train import Trainer -import utils - -if __name__ == "__main__": - - ap = argparse.ArgumentParser(description="Makes a run in trec_eval run format, given a model and a train|dev|test set" ) - ap.add_argument('model') - ap.add_argument('--word_embeddings_cache', - default='../../data/word2vec/aquaint+wiki.txt.gz.ndim=50.cache') - ap.add_argument('--paper-ext-feats', action="store_true", \ - help="external features as per the paper") - ap.add_argument('dataset_folder', help="the QA dataset folder", - choices=['../../data/TrecQA', '../../data/WikQA']) - ap.add_argument('set_split', help="train, dev or test split as the data_folder") - ap.add_argument("batch_size", help="the number of pairs to compare in each batch.\ - should be same as during training") - ap.add_argument('out_scorefile', help='output file in trec_eval format') - - - args = ap.parse_args() - - vocab_size, vec_dim = utils.load_embedding_dimensions(args.word_embeddings_cache) - - trained_model = QAModel.load(args.model) - trained_model.no_ext_feats = True - evaluator = Trainer(trained_model, 0, 0, False, vec_dim) # 0, 0, False are dummy arguments - evaluator.load_input_data(args.dataset_folder, args.word_embeddings_cache, - None, None, args.set_split, - True if args.ext_feats else False) - test_scores = evaluator.test(args.set_split, args.batch_size) - - questions, sentences, labels, vocab, maxlen_q, maxlen_s, ext_feats = \ - evaluator.data_splits[args.set_split] - - qids = [id.strip() for id in open(os.path.join(args.dataset_folder, args.set_split, 'id.txt'))\ - .readlines()] - - with open(args.out_scorefile, 'w') as outf: - old_qid = 0 - docid_c = 0 - for i in range(len(qids)): - if qids[i] != old_qid and args.dataset_folder.endswith('WikiQA'): - docid_c = 0 - old_qid = qids[i] - print('{} 0 {} 0 {} {}'.format(qids[i], docid_c, test_scores[i], - os.path.basename(args.model)), - file=outf) - docid_c += 1 diff --git a/sm_cnn/train.py b/sm_cnn/train.py index a73a457..62da27c 100644 --- a/sm_cnn/train.py +++ b/sm_cnn/train.py @@ -54,10 +54,9 @@ class Trainer(object): default_ext_feats = [np.zeros(4)] * len(self.data_splits[set_folder][0]) self.data_splits[set_folder].append(default_ext_feats) - self.embeddings[set_folder] = utils.load_cached_embeddings( \ - word_vectors_cache_file, vocab, \ - [] if "train" in set_folder else self.unk_term) - + utils.load_cached_embeddings(word_vectors_cache_file, vocab, self.embeddings, + [] if "train" in set_folder else self.unk_term) + def regularize_loss(self, loss): @@ -120,7 +119,7 @@ class Trainer(object): questions, sentences, labels, maxlen_q, maxlen_s, ext_feats = \ self.data_splits[set_folder] - word_vectors, vec_dim = self.embeddings[set_folder], self.vec_dim + word_vectors, vec_dim = self.embeddings, self.vec_dim self.model.eval() @@ -171,7 +170,7 @@ class Trainer(object): questions, sentences, labels, maxlen_q, maxlen_s, ext_feats = \ self.data_splits[set_folder] - word_vectors, vec_dim = self.embeddings[set_folder], self.vec_dim + word_vectors, vec_dim = self.embeddings, self.vec_dim # set model for training modep self.model.train() diff --git a/sm_cnn/utils.py b/sm_cnn/utils.py index 9fac336..1fb1863 100644 --- a/sm_cnn/utils.py +++ b/sm_cnn/utils.py @@ -63,7 +63,11 @@ def load_embedding_dimensions(cache_file): vocab_size, vec_dim = [int(e) for e in d.read().strip().split()] return vocab_size, vec_dim -def load_cached_embeddings(cache_file, vocab_list, oov_vec=[]): + +def load_cached_embeddings(cache_file, vocab_list, w2v_dict, oov_vec=[]): + """ + w2v_dict is filled up as reference + """ logger.debug('loading cached embeddings ') with open(cache_file + '.dimensions') as d: @@ -78,7 +82,6 @@ def load_cached_embeddings(cache_file, vocab_list, oov_vec=[]): vocab_dict = {w:k for k, w in enumerate(w2v_vocab_list)} # Read w2v for vocab appears in Q and A - w2v_dict = {} for word in vocab_list: if word in w2v_dict: continue @@ -88,10 +91,43 @@ def load_cached_embeddings(cache_file, vocab_list, oov_vec=[]): w2v_dict[word] = np.random.uniform(-0.25, 0.25, vec_dim) \ if len(oov_vec) == 0 else oov_vec #w2v_dict[word] = W[vocab_dict["unk"]] - return w2v_dict +def read_in_data(datapath, set_name, file, stop_and_stem=False, stop_punct=False, dash_split=False): + data = [] + with open(os.path.join(datapath, set_name, file)) as inf: + data = [line.strip() for line in inf.readlines()] -def read_in_dataset(dataset_folder, set_folder): + if dash_split: + def split_hyphenated_words(sentence): + rtokens = [] + for term in sentence.split(): + for t in term.split('-'): + if t: + rtokens.append(t) + return ' '.join(rtokens) + data = [split_hyphenated_words(sentence) for sentence in data] + + if stop_punct: + regex = re.compile('[{}]'.format(re.escape(string.punctuation))) + def remove_punctuation(sentence): + rtokens = [] + for term in sentence.split(): + for t in regex.sub(' ', term).strip().split(): + if t: + rtokens.append(t) + return ' '.join(rtokens) + data = [remove_punctuation(sentence) for sentence in data] + + if stop_and_stem: + stemmer = PorterStemmer() + stoplist = set(stopwords.words('english')) + def stop_stem(sentence): + return ' '.join([stemmer.stem(word) for word in sentence.split() \ + if word not in stoplist]) + data = [stop_stem(sentence) for sentence in data] + return data + +def read_in_dataset(dataset_folder, set_folder, stop_punct=False, dash_split=False): """ read in the data to return (question, sentence, label) set_folder = {train|dev|test} @@ -100,22 +136,27 @@ def read_in_dataset(dataset_folder, set_folder): max_s = 0 set_path = os.path.join(dataset_folder, set_folder) - questions = [line.strip() for line in open(os.path.join(set_path, 'a.toks')).readlines()] + #questions = [line.strip() for line in open(os.path.join(set_path, 'a.toks')).readlines()] + questions = read_in_data(dataset_folder, set_folder, "a.toks", False, stop_punct, dash_split) len_q_list = [len(q.split()) for q in questions] - sentences = [line.strip() for line in open(os.path.join(set_path, 'b.toks')).readlines()] + #sentences = [line.strip() for line in open(os.path.join(set_path, 'b.toks')).readlines()] + sentences = read_in_data(dataset_folder, set_folder, "b.toks", False, stop_punct, dash_split) len_s_list = [len(s.split()) for s in sentences] - labels = [int(line.strip()) for line in open(os.path.join(set_path, 'sim.txt')).readlines()] + #labels = [int(line.strip()) for line in open(os.path.join(set_path, 'sim.txt')).readlines()] + labels = [int(lbl) for lbl in read_in_data(dataset_folder, set_folder, "sim.txt")] - # ext_feats = [np.zeros(4)] * len(questions) - # if load_ext_features: - # ext_feats = np.array([list(map(float, line.strip().split(' '))) \ - # for line in open(os.path.join(set_path, 'overlap_feats.txt')).readlines()]) - - vocab = [line.strip() for line in open(os.path.join(dataset_folder, 'vocab.txt')).readlines()] - - return [questions, sentences, labels, max(len_q_list), max(len_s_list), vocab] + #vocab = [line.strip() for line in open(os.path.join(dataset_folder, 'vocab.txt')).readlines()] + #all_data = list(set(questions)) + list(set(sentences)) + all_data = questions + sentences + vocab_set = set() + for sentence in all_data: + for term in sentence.split(): + vocab_set.add(term) + vocab = list(vocab_set) + + return [questions, sentences, labels, max(len_q_list), max(len_s_list), vocab] def get_test_qids_labels(dataset_folder, set_folder): @@ -129,7 +170,8 @@ if __name__ == "__main__": vocab = ["unk", "idontreallythinkthiswordexists", "hello"] - w2v_dict, vec_dim = load_cached_embeddings("../../data/word2vec/aquaint+wiki.txt.gz.ndim=50.cache", vocab) + w2v_dict = {} + load_cached_embeddings("../../data/word2vec/aquaint+wiki.txt.gz.ndim=50.cache", vocab, w2v_dict) for w, v in w2v_dict.iteritems(): print(w)