diff --git a/setup.py b/setup.py index d60726d..f052898 100644 --- a/setup.py +++ b/setup.py @@ -3,5 +3,5 @@ from setuptools import setup setup(name='castor', version='1.0.0', description='models for question answering', - packages=['sm_model'], + packages=['sm_cnn'], ) diff --git a/sm_model/.gitignore b/sm_cnn/.gitignore similarity index 100% rename from sm_model/.gitignore rename to sm_cnn/.gitignore diff --git a/sm_model/README.md b/sm_cnn/README.md similarity index 84% rename from sm_model/README.md rename to sm_cnn/README.md index 2c6ea35..b63d99d 100644 --- a/sm_model/README.md +++ b/sm_cnn/README.md @@ -28,16 +28,16 @@ git clone https://github.com/castorini/Castor.git This should generate: ``` ├── Castor -│   ├── castorini_smmodel_bridge.py -│   ├── README.md -│   └── sm_model/ +│   ├── idf_baseline +│   ├── kim_cnn +│   └── sm_cnn ├── data │   ├── README.md │   ├── TrecQA/ │   └── word2vec/ └── models ├── README.md - └── sm_model/ + └── sm_cnn/ ``` 2. Preprocess data @@ -57,7 +57,7 @@ python3 build_vocab.py Make trec_eval ``` -cd Castor/sm_model/ +cd Castor/sm_cnn/ cd trec_eval-8.0 make clean && make cd .. @@ -65,9 +65,9 @@ cd .. To train the S&M model on TrecQA ``` -python main.py ../../model/sm_model/sm_model.train-all +python main.py ../../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor ``` -The final model will be saved to ```../../model/sm_model/sm_model.train-all``` +The final model will be saved to ```../../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor``` _NOTE:_ On first run, the program will create a memory-mapped cache for word e mbeddings (943MB) in ``data/word2vec``. The cache allows for faster loading of data in future runs. diff --git a/sm_model/__init__.py b/sm_cnn/__init__.py similarity index 100% rename from sm_model/__init__.py rename to sm_cnn/__init__.py diff --git a/sm_cnn/bridge.py b/sm_cnn/bridge.py new file mode 100644 index 0000000..07d88d8 --- /dev/null +++ b/sm_cnn/bridge.py @@ -0,0 +1,186 @@ +import json +import os +import sys +from collections import Counter +import argparse + +import numpy as np +import torch +from nltk.tokenize import TreebankWordTokenizer +from torch.autograd import Variable +from py4j.java_gateway import JavaGateway + +from sm_cnn import model +from sm_cnn.external_features import compute_overlap, compute_idf_weighted_overlap, stopped + +sys.modules['model'] = model + + +class SMModelBridge(object): + + def __init__(self, model_file, word_embeddings_cache_file, index_path): + # init torch random seeds + torch.manual_seed(1234) + np.random.seed(1234) + + # load model + self.model = model.QAModel.load(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) + + self.index = index_path + + + def _preload_cached_embeddings(self, cache_file): + + with open(cache_file + '.dimensions') as d: + vocab_size, vec_dim = [int(e) for e in d.read().strip().split()] + + self.W = np.memmap(cache_file, dtype=np.double, shape=(vocab_size, vec_dim)) + + with open(cache_file + '.vocab') as f: + w2v_vocab_list = map(str.strip, f.readlines()) + + self.vocab_dict = {w:k for k, w in enumerate(w2v_vocab_list)} + return vec_dim + + + def parse(self, sentence): + s_toks = TreebankWordTokenizer().tokenize(sentence) + s_str = ' '.join(s_toks).lower() + return s_str + + + def make_input_matrix(self, sentence): + terms = sentence.strip().split() + # word_embeddings = torch.zeros(max_len, vec_dim).type(torch.DoubleTensor) + word_embeddings = torch.zeros(len(terms), self.vec_dim).type(torch.DoubleTensor) + for i in range(len(terms)): + word = terms[i] + if word not in self.vocab_dict: + emb = torch.from_numpy(self.unk_term_vec) + else: + emb = torch.from_numpy(self.W[self.vocab_dict[word]]) + word_embeddings[i] = emb + input_tensor = torch.zeros(1, self.vec_dim, len(terms)) + input_tensor[0] = torch.transpose(word_embeddings, 0, 1) + return input_tensor + + + def get_tensorized_inputs(self, batch_ques, batch_sents, batch_ext_feats): + assert(1 == len(batch_ques)) + tensorized_inputs = [] + for i in range(len(batch_ques)): + xq = Variable(self.make_input_matrix(batch_ques[i])) + xs = Variable(self.make_input_matrix(batch_sents[i])) + ext_feats = Variable(torch.FloatTensor(batch_ext_feats[i])) + ext_feats = torch.unsqueeze(ext_feats, 0) + tensorized_inputs.append((xq, xs, ext_feats)) + return tensorized_inputs + + def rerank_candidate_answers(self, question, answers, idf_json): + # run through the model + scores_sentences = [] + question = self.parse(question) + term_idfs = json.loads(idf_json) + term_idfs = dict((k, float(v)) for k, v in term_idfs.items()) + + for answer in answers: + answer = self.parse(answer) + + overlap = compute_overlap([question], [answer]) + idf_weighted_overlap = compute_idf_weighted_overlap([question], [answer], term_idfs) + overlap_no_stopwords =\ + compute_overlap(stopped([question]), stopped([answer])) + idf_weighted_overlap_no_stopwords =\ + compute_idf_weighted_overlap(stopped([question]), stopped([answer]), term_idfs) + ext_feats = [np.array(feats) for feats in zip(overlap, idf_weighted_overlap,\ + overlap_no_stopwords, idf_weighted_overlap_no_stopwords)] + + xq, xa, x_ext_feats = self.get_tensorized_inputs([question], [answer], \ + ext_feats)[0] + pred = self.model(xq, xa, x_ext_feats) + pred = torch.exp(pred) + scores_sentences.append((pred.data.squeeze()[1], answer)) + + return scores_sentences + +def get_term_idf_json_list(index_path, sent_list): + gateway = JavaGateway() + index = gateway.jvm.java.lang.String(index_path) + pyserini = gateway.jvm.io.anserini.py4j.PyseriniEntryPoint() + pyserini.initializeWithIndex(index_path) + java_list = gateway.jvm.java.util.ArrayList() + + for l in sent_list: + java_list.add(l) + + json_object = pyserini.getTermIdfJSONs(java_list) + return json_object + +if __name__ == "__main__": + ap = argparse.ArgumentParser(description="Bridge Demo. Produces scores in trec_eval format", + formatter_class=argparse.ArgumentDefaultsHelpFormatter) + ap.add_argument('model', help="the path to the saved model file") + ap.add_argument('--word-embeddings-cache', help="the embeddings 'cache' file",\ + default='../data/word2vec/aquaint+wiki.txt.gz.ndim=50.cache') + ap.add_argument('index_path', help="the path to the source corpus index") + # 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 {TrecQA|WikiQA}", + default='../data/TrecQA/') + + args = ap.parse_args() + + smmodel = SMModelBridge( + args.model, + args.word_embeddings_cache, + args.index_path + ) + + train_set, dev_set, test_set = 'train', 'dev', 'test' + if 'TrecQA' in args.dataset_folder: + train_set, dev_set, test_set = 'train-all', 'raw-dev', 'raw-test' + + + + for split in [dev_set, test_set]: + outfile = open('bridge.{}.scores'.format(split), 'w') + + questions = [q.strip() for q in \ + open(os.path.join(args.dataset_folder, split, 'a.toks')).readlines()] + answers = [q.strip() for q in \ + open(os.path.join(args.dataset_folder, split, 'b.toks')).readlines()] + labels = [q.strip() for q in \ + open(os.path.join(args.dataset_folder, split, 'sim.txt')).readlines()] + qids = [q.strip() for q in \ + open(os.path.join(args.dataset_folder, split, 'id.txt')).readlines()] + + qid_question = dict(zip(qids, questions)) + q_counts = Counter(questions) + + answers_offset = 0 + docid_counter = 0 + + all_questions_answers = questions + answers + idf_json = get_term_idf_json_list(args.index_path, all_questions_answers) + + for qid, question in sorted(qid_question.items(), key=lambda x: float(x[0])): + num_answers = q_counts[question] + q_answers = answers[answers_offset: answers_offset + num_answers] + answers_offset += num_answers + sentence_scores = smmodel.rerank_candidate_answers(question, q_answers, idf_json) + + for score, sentence in sentence_scores: + print('{} Q0 {} 0 {} sm_cnn_bridge.{}.run'.format( + qid, + docid_counter, + score, + os.path.basename(args.dataset_folder) + ), file=outfile) + docid_counter += 1 + if 'WikiQA' in args.dataset_folder: + docid_counter = 0 + + outfile.close() diff --git a/sm_model/external_features.py b/sm_cnn/external_features.py similarity index 100% rename from sm_model/external_features.py rename to sm_cnn/external_features.py diff --git a/sm_model/main.py b/sm_cnn/main.py similarity index 100% rename from sm_model/main.py rename to sm_cnn/main.py diff --git a/sm_model/make_run.py b/sm_cnn/make_run.py similarity index 100% rename from sm_model/make_run.py rename to sm_cnn/make_run.py diff --git a/sm_model/model.py b/sm_cnn/model.py similarity index 100% rename from sm_model/model.py rename to sm_cnn/model.py diff --git a/sm_model/requirements.txt b/sm_cnn/requirements.txt similarity index 100% rename from sm_model/requirements.txt rename to sm_cnn/requirements.txt diff --git a/sm_model/run_eval.sh b/sm_cnn/run_eval.sh similarity index 100% rename from sm_model/run_eval.sh rename to sm_cnn/run_eval.sh diff --git a/sm_model/train.py b/sm_cnn/train.py similarity index 100% rename from sm_model/train.py rename to sm_cnn/train.py diff --git a/sm_model/trec_eval-8.0/Makefile b/sm_cnn/trec_eval-8.0/Makefile similarity index 100% rename from sm_model/trec_eval-8.0/Makefile rename to sm_cnn/trec_eval-8.0/Makefile diff --git a/sm_model/trec_eval-8.0/README b/sm_cnn/trec_eval-8.0/README similarity index 100% rename from sm_model/trec_eval-8.0/README rename to sm_cnn/trec_eval-8.0/README diff --git a/sm_model/trec_eval-8.0/buf.h b/sm_cnn/trec_eval-8.0/buf.h similarity index 100% rename from sm_model/trec_eval-8.0/buf.h rename to sm_cnn/trec_eval-8.0/buf.h diff --git a/sm_model/trec_eval-8.0/buf_util.c b/sm_cnn/trec_eval-8.0/buf_util.c similarity index 100% rename from sm_model/trec_eval-8.0/buf_util.c rename to sm_cnn/trec_eval-8.0/buf_util.c diff --git a/sm_model/trec_eval-8.0/common.h b/sm_cnn/trec_eval-8.0/common.h similarity index 100% rename from sm_model/trec_eval-8.0/common.h rename to sm_cnn/trec_eval-8.0/common.h diff --git a/sm_model/trec_eval-8.0/error_msgs.c b/sm_cnn/trec_eval-8.0/error_msgs.c similarity index 100% rename from sm_model/trec_eval-8.0/error_msgs.c rename to sm_cnn/trec_eval-8.0/error_msgs.c diff --git a/sm_model/trec_eval-8.0/form_trvec.c b/sm_cnn/trec_eval-8.0/form_trvec.c similarity index 100% rename from sm_model/trec_eval-8.0/form_trvec.c rename to sm_cnn/trec_eval-8.0/form_trvec.c diff --git a/sm_model/trec_eval-8.0/get_qrels.c b/sm_cnn/trec_eval-8.0/get_qrels.c similarity index 100% rename from sm_model/trec_eval-8.0/get_qrels.c rename to sm_cnn/trec_eval-8.0/get_qrels.c diff --git a/sm_model/trec_eval-8.0/get_top.c b/sm_cnn/trec_eval-8.0/get_top.c similarity index 100% rename from sm_model/trec_eval-8.0/get_top.c rename to sm_cnn/trec_eval-8.0/get_top.c diff --git a/sm_model/trec_eval-8.0/smart_error.h b/sm_cnn/trec_eval-8.0/smart_error.h similarity index 100% rename from sm_model/trec_eval-8.0/smart_error.h rename to sm_cnn/trec_eval-8.0/smart_error.h diff --git a/sm_model/trec_eval-8.0/sysfunc.h b/sm_cnn/trec_eval-8.0/sysfunc.h similarity index 100% rename from sm_model/trec_eval-8.0/sysfunc.h rename to sm_cnn/trec_eval-8.0/sysfunc.h diff --git a/sm_model/trec_eval-8.0/tr_vec.h b/sm_cnn/trec_eval-8.0/tr_vec.h similarity index 100% rename from sm_model/trec_eval-8.0/tr_vec.h rename to sm_cnn/trec_eval-8.0/tr_vec.h diff --git a/sm_model/trec_eval-8.0/trec_eval.c b/sm_cnn/trec_eval-8.0/trec_eval.c similarity index 100% rename from sm_model/trec_eval-8.0/trec_eval.c rename to sm_cnn/trec_eval-8.0/trec_eval.c diff --git a/sm_model/utils.py b/sm_cnn/utils.py similarity index 100% rename from sm_model/utils.py rename to sm_cnn/utils.py diff --git a/sm_model/bridge.py b/sm_model/bridge.py deleted file mode 100644 index 00aa771..0000000 --- a/sm_model/bridge.py +++ /dev/null @@ -1,202 +0,0 @@ -import os -import sys -import pickle -import string -from collections import defaultdict - -import numpy as np -import torch -from nltk.tokenize import TreebankWordTokenizer -from torch.autograd import Variable - -from sm_model import model - -sys.modules['model'] = model - - -class SMModelBridge(object): - - def __init__(self, model_file, word_embeddings_cache_file, stopwords_file, word2dfs_file): - # init torch random seeds - torch.manual_seed(1234) - np.random.seed(1234) - - # load model - self.model = model.QAModel.load(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) - - # stopwords - self.stoplist = set([line.strip() for line in open(stopwords_file)]) - - # word dfs - if os.path.isfile(word2dfs_file): - with open(word2dfs_file, "rb") as w2dfin: - self.word2dfs = pickle.load(w2dfin) - - - def _preload_cached_embeddings(self, cache_file): - - with open(cache_file + '.dimensions') as d: - vocab_size, vec_dim = [int(e) for e in d.read().strip().split()] - - self.W = np.memmap(cache_file, dtype=np.double, shape=(vocab_size, vec_dim)) - - with open(cache_file + '.vocab') as f: - w2v_vocab_list = map(str.strip, f.readlines()) - - self.vocab_dict = {w:k for k, w in enumerate(w2v_vocab_list)} - return vec_dim - - - def parser(self, q, a): - q_toks = TreebankWordTokenizer().tokenize(q) - q_str = ' '.join(q_toks).lower() - a_list = [] - for ans in a: - ans_toks = TreebankWordTokenizer().tokenize(ans) - a_str = ' '.join(ans_toks).lower() - a_list.append(a_str) - return q_str, a_list - - - def compute_overlap_features(self, q_str, a_list, word2df=None, stoplist=None): - word2df = word2df if word2df else {} - stoplist = stoplist if stoplist else set() - feats_overlap = [] - for a in a_list: - question = q_str.split() - answer = a.split() - # q_set = set(question) - # a_set = set(answer) - q_set = set([q for q in question if q not in stoplist]) - a_set = set([a for a in answer if a not in stoplist]) - word_overlap = q_set.intersection(a_set) - # overlap = float(len(word_overlap)) / (len(q_set) * len(a_set) + 1e-8) - if len(q_set) == 0 and len(a_set) == 0: - overlap = 0 - else: - overlap = float(len(word_overlap)) / (len(q_set) + len(a_set)) - - # q_set = set([q for q in question if q not in stoplist]) - # a_set = set([a for a in answer if a not in stoplist]) - word_overlap = q_set.intersection(a_set) - df_overlap = 0.0 - for w in word_overlap: - df_overlap += word2df[w] - - if len(q_set) == 0 and len(a_set) == 0: - df_overlap = 0 - else: - df_overlap /= (len(q_set) + len(a_set)) - - feats_overlap.append(np.array([overlap, df_overlap])) - return np.array(feats_overlap) - - - def make_input_matrix(self, sentence): - terms = sentence.strip().split() - # word_embeddings = torch.zeros(max_len, vec_dim).type(torch.DoubleTensor) - word_embeddings = torch.zeros(len(terms), self.vec_dim).type(torch.DoubleTensor) - for i in range(len(terms)): - word = terms[i] - if word not in self.vocab_dict: - emb = torch.from_numpy(self.unk_term_vec) - else: - emb = torch.from_numpy(self.W[self.vocab_dict[word]]) - word_embeddings[i] = emb - input_tensor = torch.zeros(1, self.vec_dim, len(terms)) - input_tensor[0] = torch.transpose(word_embeddings, 0, 1) - return input_tensor - - - def get_tensorized_inputs(self, batch_ques, batch_sents, batch_ext_feats): - assert(1 == len(batch_ques)) - tensorized_inputs = [] - for i in range(len(batch_ques)): - xq = Variable(self.make_input_matrix(batch_ques[i])) - xs = Variable(self.make_input_matrix(batch_sents[i])) - ext_feats = Variable(torch.FloatTensor(batch_ext_feats[i])) - ext_feats = torch.unsqueeze(ext_feats, 0) - tensorized_inputs.append((xq, xs, ext_feats)) - return tensorized_inputs - - - def rerank_candidate_answers(self, question, answers): - # tokenize - q_str, a_list = self.parser(question, answers) - - # calculate overlap features - overlap_feats = self.compute_overlap_features(q_str, a_list, \ - stoplist=None, word2df=self.word2dfs) - overlap_feats_stoplist = self.compute_overlap_features(q_str, a_list, \ - stoplist=self.stoplist, word2df=self.word2dfs) - overlap_feats_vec = np.hstack([overlap_feats, overlap_feats_stoplist]) - - # run through the model - scores_sentences = [] - for i in range(len(a_list)): - xq, xa, x_ext_feats = self.get_tensorized_inputs([q_str], [a_list[i]], \ - [overlap_feats_vec[i]])[0] - pred = self.model(xq, xa, x_ext_feats) - pred = torch.exp(pred) - scores_sentences.append((pred.data.squeeze()[1], a_list[i])) - - return scores_sentences - - -if __name__ == "__main__": - - ap = argparse.ArgumentParser(description="Bridge Demo. Produces scores in trec_eval format") - ap.add_argument('model') - ap.add_argument('--word_embeddings_cache', default='../data/word2vec/aquaint+wiki.txt.gz.ndim=50.cache') - ap.add_argument('--stopwords_file', default='../data/TrecQA/stopwords.txt') - ap.add_argument('--wordDF_file', default='../data/TrecQA/word2dfs.p') - ap.add_argument('--no_ext_feats', action="store_true", help="This argument has no effect because the model saves its members") - ap.add_argument('--use_pre_ext_feats', action="store_true", help="use the precomputed external overlap features") - ap.add_argument('--data_folder', default='../data/TrecQA/') - ap.add_argument('dataset', choices=['train-all', 'raw-test', 'raw-dev', 'train']) - ap.add_argument('out_scorefile', help='file in trec_eval format') - ap.add_argument('--out_qrels', help='will also output qrels trec_eval format') - - args = ap.parse_args() - - smmodel = SMModelBridge( - #'../models/sm_model/sm_model.TrecQA.TRAIN-ALL.2017-04-02.castor', - args.model, - args.word_embeddings_cache, - args.stopwords_file, - args.wordDF_file) - - # if args.no_ext_feats: - # smmodel.model.no_ext_feats = True - - - allque = [q.strip() for q in open(os.path.join('../data/TrecQA/', args.dataset+'/a.toks')).readlines()] - allans = [a.strip() for a in open(os.path.join('../data/TrecQA/', args.dataset+'/b.toks')).readlines()] - labels = [y.strip() for y in open(os.path.join('../data/TrecQA/', args.dataset+'/sim.txt')).readlines()] - qids = [id.strip() for id in open(os.path.join('../data/TrecQA/', args.dataset+'/id.txt')).readlines()] - - pre_ext_feats = None - if args.use_pre_ext_feats: - pre_ext_feats = [ [float(e) for e in x.split() ] for x in open(os.path.join('../data/TrecQA/', args.dataset+'/overlap_feats.txt')).readlines()] - - scoref = open(args.out_scorefile, 'w') - if args.out_qrels: - qrelf = open(args.out_qrels, 'w') - - for i in range(len(allque)): - question = allque[i] - answers = [allans[i]] - ext_feats = None - if args.use_pre_ext_feats: - ext_feats = [pre_ext_feats[i]] - ss = smmodel.rerank_candidate_answers(question, answers, ext_feats) - # print('Question:', question) - for score, sentence in ss: - #print(score, '\t', sentence) - #print('{}\t{}'.format(labels[i], score)) - print('{} {} {} {} {} {}'.format(qids[i], '0', i, 0, score, 'sm_model.'+args.dataset), file=scoref) - if args.out_qrels: - print('{} {} {} {}'.format(qids[i], '0', i, labels[i]), file=qrelf)