From 36b1ddb8588145dd91d338e28d7d0ca504c9ee31 Mon Sep 17 00:00:00 2001 From: rosequ Date: Wed, 4 Oct 2017 15:59:26 -0400 Subject: [PATCH] SM model for WikiQA (#63) (#64) * support for WikiQA dataset * parallel runs for both datasets * minor fixes * updated README * removed data folder; added scripts to create dataset; updated README * after CR * after CR2 --- sm_modified_cnn/.gitignore | 1 + sm_modified_cnn/README.md | 68 ++++++++++-- sm_modified_cnn/args.py | 2 +- sm_modified_cnn/create_dataset.sh | 17 +++ sm_modified_cnn/evaluate.py | 9 +- sm_modified_cnn/main.py | 17 ++- sm_modified_cnn/overlap_features.py | 154 ++++++++++++++++++++++++++++ sm_modified_cnn/train.py | 12 ++- sm_modified_cnn/wiki_dataset.py | 16 +++ 9 files changed, 277 insertions(+), 19 deletions(-) create mode 100755 sm_modified_cnn/create_dataset.sh create mode 100644 sm_modified_cnn/overlap_features.py create mode 100644 sm_modified_cnn/wiki_dataset.py diff --git a/sm_modified_cnn/.gitignore b/sm_modified_cnn/.gitignore index 411a968..7483f8e 100644 --- a/sm_modified_cnn/.gitignore +++ b/sm_modified_cnn/.gitignore @@ -3,3 +3,4 @@ text/ trained_models/ trec_eval-8.0/trec_eval.dSYM +data/ diff --git a/sm_modified_cnn/README.md b/sm_modified_cnn/README.md index 95f573b..173e5ec 100644 --- a/sm_modified_cnn/README.md +++ b/sm_modified_cnn/README.md @@ -29,11 +29,51 @@ make cd .. ``` + +### Setup +Clone and create the dataset: +```bash +git clone https://github.com/castorini/data.git +git clone https://github.com/castorini/Castor.git +``` + +You should you see the following tree: +``` +. +├── Castor +│   ├── README.md +│   ├── baseline_results.tsv +│   ├── idf_baseline +│   ├── kim_cnn +│   ├── mp_cnn +│   ├── setup.py +│   ├── sm_cnn +│   └── sm_modified_cnn +└── data + ├── GloVe + ├── ParagramEmbeddings + ├── README.md + ├── SimpleQuestions_v2 + ├── TrecQA + ├── WikiQA + ├── msrvid + ├── requirements.txt + ├── sick + ├── twitterPPDB + ├── utils + └── word2vec +``` + +To create the dataset: +```bash +cd Castor/sm_modified_cnn/ +./create_dataset.sh +``` + +### Training Download the word2vec model from [here] (https://drive.google.com/file/d/0B2u_nClt6NbzUmhOZU55eEo4QWM/view?usp=sharing) and copy it to the `data/` folder. -### Training the model - You can train the SM model for the 4 following configurations: 1. __random__ - the word embedddings are initialized randomly and are tuned during training 2. __static__ - the word embeddings are static (Severyn and Moschitti, SIGIR'15) @@ -63,16 +103,32 @@ python main.py --trained_model saves/TREC/multichannel_best_model.pt The performance on TrecQA dataset: -### Best dev +### TrecQA: +#### Best dev Metric |rand |static|non-static|multichannel -------|------|------|----------|------------ MAP |0.8096|0.8162|0.8387 | 0.8274 MRR |0.8560|0.8918|0.9058 | 0.8818 -### Test - +#### Test Metric |rand |static|non-static|multichannel -------|-------|------|----------|------------ MAP |0.7441 |0.7524|0.7688 |0.7641 -MRR |0.8172 |0.8012|0.8144 |0.8174 \ No newline at end of file +MRR |0.8172 |0.8012|0.8144 |0.8174 + +### WikiQA: + +#### Best dev +Metric |rand |static|non-static|multichannel +-------|------|------|----------|------------ +MAP |0.7109|0.7204|0.7049 | 0.7245 +MRR |0.7169|0.7234|0.7075 | 0.7259 + +#### Test +Metric |rand |static|non-static|multichannel +-------|-------|------|----------|------------ +MAP |0.6313 |0.6378|0.6455 |0.6476 +MRR |0.6522 |0.6542|0.6689 |0.6646 + +NB: The results on WikiQA are based on the SM model hyperparameters. \ No newline at end of file diff --git a/sm_modified_cnn/args.py b/sm_modified_cnn/args.py index 038105b..3d75b58 100644 --- a/sm_modified_cnn/args.py +++ b/sm_modified_cnn/args.py @@ -9,7 +9,7 @@ def get_args(): parser.add_argument('--mode', type=str, default='static') parser.add_argument('--lr', type=float, default=1.0) parser.add_argument('--seed', type=int, default=3435) - parser.add_argument('--dataset', type=str, default='TREC') + parser.add_argument('--dataset', type=str, help='TREC|wiki', default='TREC') parser.add_argument('--resume_snapshot', type=str, default=None) parser.add_argument('--dev_every', type=int, default=30) parser.add_argument('--log_every', type=int, default=10) diff --git a/sm_modified_cnn/create_dataset.sh b/sm_modified_cnn/create_dataset.sh new file mode 100755 index 0000000..e8d7903 --- /dev/null +++ b/sm_modified_cnn/create_dataset.sh @@ -0,0 +1,17 @@ +#!/bin/sh +mkdir -p data +python overlap_features.py --dir ../../data/TrecQA/ + +CURRENT_DIR=$(pwd) +cd ../../data/TrecQA +cd raw-dev/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/trecqa.dev.tsv; cd .. +cd raw-test/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/trecqa.test.tsv; cd .. +cd train-all/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/trecqa.train.tsv; cd .. +cd $CURRENT_DIR + +python overlap_features.py --dir ../../data/WikiQA/ +cd ../../data/WikiQA +cd dev/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/wikiqa.dev.tsv; cd .. +cd test/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/wikiqa.test.tsv; cd .. +cd train/; paste id.txt sim.txt a.toks b.toks overlap_feats.txt > $CURRENT_DIR/data/wikiqa.train.tsv; cd .. +cd $CURRENT_DIR \ No newline at end of file diff --git a/sm_modified_cnn/evaluate.py b/sm_modified_cnn/evaluate.py index 4460dd8..76ba31d 100644 --- a/sm_modified_cnn/evaluate.py +++ b/sm_modified_cnn/evaluate.py @@ -1,9 +1,10 @@ import shlex import subprocess -def evaluate(instances, valid, config): +def evaluate(instances, dataset, valid, config): sorted_instances = sorted(instances, key=lambda x: (x[0])) - with open('{}.{}.run.txt'.format(valid, config), 'w') as run, open('{}.{}.qrel.txt'.format(valid, config), 'w') as qrel: + with open('{}.{}.{}.run.txt'.format(dataset, valid, config), 'w') as run, \ + open('{}.{}.{}.qrel.txt'.format(dataset, valid, config), 'w') as qrel: i = 0 for instance in sorted_instances: qid, predicted, score, gold = instance[0], instance[1], instance[2], instance[3] @@ -13,8 +14,8 @@ def evaluate(instances, valid, config): qrel.write('{} 0 {} {}\n'.format(qid, i, gold)) i += 1 - pargs = shlex.split("./eval/trec_eval.9.0/trec_eval -m map -m recip_rank {}.{}.qrel.txt {}.{}.run.txt" - .format(valid, config, valid, config)) + pargs = shlex.split("./eval/trec_eval.9.0/trec_eval -m map -m recip_rank {}.{}.{}.qrel.txt {}.{}.{}.run.txt" + .format(dataset, valid, config, dataset, valid, config)) p = subprocess.Popen(pargs, stdout=subprocess.PIPE, stderr=subprocess.PIPE) pout, perr = p.communicate() lines = pout.split(b'\n') diff --git a/sm_modified_cnn/main.py b/sm_modified_cnn/main.py index 50e0d0b..a7a97df 100644 --- a/sm_modified_cnn/main.py +++ b/sm_modified_cnn/main.py @@ -7,6 +7,7 @@ from torchtext import data from args import get_args from trec_dataset import TrecDataset +from wiki_dataset import WikiDataset from evaluate import evaluate logger = logging.getLogger(__name__) @@ -41,7 +42,13 @@ LABEL = data.Field(sequential=False) EXTERNAL = data.Field(sequential=False, tensor_type=torch.FloatTensor, batch_first=True, use_vocab=False, preprocessing=data.Pipeline(lambda x: x.split()), postprocessing=data.Pipeline(lambda x, train: [float(y) for y in x])) -train, dev, test = TrecDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +if config.dataset == 'trec': + train, dev, test = TrecDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +elif config.dataset == 'wiki': + train, dev, test = WikiDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +else: + print("Unsupported dataset") + exit() QID.build_vocab(train, dev, test) QUESTION.build_vocab(train, dev, test) @@ -68,7 +75,7 @@ else: index2label = np.array(LABEL.vocab.itos) index2qid = np.array(QID.vocab.itos) -def predict(test_mode, dataset_iter): +def predict(dataset, test_mode, dataset_iter): model.eval() dataset_iter.init_epoch() @@ -88,11 +95,11 @@ def predict(test_mode, dataset_iter): true_label_array[i] instance.append((this_qid, predicted_label, score, gold_label)) - dev_map, dev_mrr = evaluate(instance, test_mode, config.mode) + dev_map, dev_mrr = evaluate(instance, dataset, test_mode, config.mode) print(dev_map, dev_mrr) # Run the model on the dev set -predict('dev', dataset_iter=dev_iter) +predict(config.dataset, 'dev', dataset_iter=dev_iter) # Run the model on the test set -predict('test', dataset_iter=test_iter) +predict(config.dataset, 'test', dataset_iter=test_iter) diff --git a/sm_modified_cnn/overlap_features.py b/sm_modified_cnn/overlap_features.py new file mode 100644 index 0000000..9a799e3 --- /dev/null +++ b/sm_modified_cnn/overlap_features.py @@ -0,0 +1,154 @@ +import numpy as np +import string +import pickle +from collections import defaultdict +from argparse import ArgumentParser + +from nltk.stem.porter import PorterStemmer + +def load_data(dname): + stemmer = PorterStemmer() + qids, questions, answers, labels = [], [], [], [] + print('Load folder ' + dname) + with open(dname+'a.toks', encoding='utf-8') as f: + for line in f: + question = line.strip().split() + question = [stemmer.stem(word) for word in question] + questions.append(question) + with open(dname+'b.toks', encoding='utf-8') as f: + for line in f: + answer = line.strip().split() + answer_list = [] + for word in answer: + try: + answer_list.append(stemmer.stem(word)) + except Exception as e: + print("couldn't stem the word:" + word) + answers.append(answer_list) + with open(dname+'id.txt', encoding='utf-8') as f: + for line in f: + qids.append(line.strip()) + with open(dname+'sim.txt', encoding='utf-8') as f: + for line in f: + labels.append(int(line.strip())) + return qids, questions, answers, labels + +def compute_overlap_features(questions, answers, word2df=None, stoplist=None): + word2df = word2df if word2df else {} + stoplist = stoplist if stoplist else set() + feats_overlap = [] + for question, answer in zip(questions, answers): + 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) + if len(q_set) == 0 and len(a_set) == 0: + overlap = 0 + else: + overlap = float(len(word_overlap)) / (len(q_set) + len(a_set)) + + 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 compute_overlap_idx(questions, answers, stoplist, q_max_sent_length, a_max_sent_length): + stoplist = stoplist if stoplist else [] + q_indices, a_indices = [], [] + for question, answer in zip(questions, answers): + 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) + + q_idx = np.ones(q_max_sent_length) * 2 + for i, q in enumerate(question): + value = 0 + if q in word_overlap: + value = 1 + q_idx[i] = value + q_indices.append(q_idx) + + a_idx = np.ones(a_max_sent_length) * 2 + for i, a in enumerate(answer): + value = 0 + if a in word_overlap: + value = 1 + a_idx[i] = value + a_indices.append(a_idx) + + q_indices = np.vstack(q_indices).astype('int32') + a_indices = np.vstack(a_indices).astype('int32') + + return q_indices, a_indices + +def compute_dfs(docs): + word2df = defaultdict(float) + for doc in docs: + for w in set(doc): + word2df[w] += 1.0 + num_docs = len(docs) + + for w, value in word2df.items(): + word2df[w] = np.math.log(num_docs / value) # bug feats fixed + + return word2df + +if __name__ == '__main__': + parser = ArgumentParser(description='create TrecQA/WikiQA dataset') + parser.add_argument('--dir', help='path to the TrecQA|WikiQA data directory', default="../../data/TrecQA") + args = parser.parse_args() + + stoplist = set([line.strip() for line in open('../../data/TrecQA/stopwords.txt', encoding='utf-8')]) + punct = set(string.punctuation) + stoplist.update(punct) + + all_questions, all_answers, all_qids = [], [], [] + base_dir = args.dir + + if 'TrecQA' in base_dir: + sub_dirs = ['train/', 'train-all/', 'raw-dev/', 'raw-test/'] + elif 'WikiQA' in base_dir: + sub_dirs = ['train/', 'dev/', 'test/'] + else: + print('Unsupported dataset') + exit() + + for sub in sub_dirs: + qids, questions, answers, labels = load_data(base_dir + sub) + all_questions.extend(questions) + all_answers.extend(answers) + all_qids.extend(qids) + + seen = set() + unique_questions = [] + for q, qid in zip(all_questions, all_qids): + if qid not in seen: + seen.add(qid) + unique_questions.append(q) + + docs = all_answers + unique_questions + word2dfs = compute_dfs(docs) + pickle.dump(word2dfs, open("word2dfs.p", "wb")) + + q_max_sent_length = max(map(lambda x: len(x), all_questions)) + a_max_sent_length = max(map(lambda x: len(x), all_answers)) + + for sub in sub_dirs: + qids, questions, answers, labels = load_data(base_dir + sub) + + overlap_feats = compute_overlap_features(questions, answers, stoplist=None, word2df=word2dfs) + overlap_feats_stoplist = compute_overlap_features(questions, answers, stoplist=stoplist, word2df=word2dfs) + overlap_feats = np.hstack([overlap_feats, overlap_feats_stoplist]) + + with open(base_dir + sub + 'overlap_feats.txt', 'w') as f: + for i in range(overlap_feats.shape[0]): + for j in range(4): + f.write(str(overlap_feats[i][j]) + ' ') + f.write('\n') diff --git a/sm_modified_cnn/train.py b/sm_modified_cnn/train.py index 367b159..1f6a601 100644 --- a/sm_modified_cnn/train.py +++ b/sm_modified_cnn/train.py @@ -3,7 +3,6 @@ import os import numpy as np import random -import logging import torch import torch.nn as nn from torchtext import data @@ -11,6 +10,7 @@ from torchtext import data from args import get_args from model import SmPlusPlus from trec_dataset import TrecDataset +from wiki_dataset import WikiDataset from evaluate import evaluate args = get_args() @@ -73,7 +73,13 @@ LABEL = data.Field(sequential=False) EXTERNAL = data.Field(sequential=False, tensor_type=torch.FloatTensor, batch_first=True, use_vocab=False, preprocessing=data.Pipeline(lambda x: x.split()), postprocessing=data.Pipeline(lambda x, train: [float(y) for y in x])) -train, dev, test = TrecDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +if config.dataset == 'TREC': + train, dev, test = TrecDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +elif config.dataset == 'wiki': + train, dev, test = WikiDataset.splits(QID, QUESTION, ANSWER, EXTERNAL, LABEL) +else: + print("Unsupported dataset") + exit() QID.build_vocab(train, dev, test) QUESTION.build_vocab(train, dev, test) @@ -188,7 +194,7 @@ while True: instance.append((this_qid, predicted_label, score, gold_label)) - dev_map, dev_mrr = evaluate(instance, 'valid', config.mode) + dev_map, dev_mrr = evaluate(instance, config.dataset, 'valid', config.mode) print(dev_log_template.format(time.time() - start, epoch, iterations, 1 + batch_idx, len(train_iter), 100. * (1 + batch_idx) / len(train_iter), loss.data[0], diff --git a/sm_modified_cnn/wiki_dataset.py b/sm_modified_cnn/wiki_dataset.py new file mode 100644 index 0000000..67cbbf9 --- /dev/null +++ b/sm_modified_cnn/wiki_dataset.py @@ -0,0 +1,16 @@ +from torchtext import data +import os + +class WikiDataset(data.TabularDataset): + dirname = 'data' + @classmethod + + def splits(cls, question_id, question_field, answer_field, external_field, label_field, + train='train.tsv', validation='dev.tsv', test='test.tsv'): + path = './data' + prefix_name = 'wikiqa.' + return super(WikiDataset, cls).splits( + os.path.join(path, prefix_name), train, validation, test, + format='TSV', fields=[('qid', question_id), ('label', label_field), ('question', question_field), + ('answer', answer_field), ('ext_feat', external_field)] + )