mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
* 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
This commit is contained in:
@@ -3,3 +3,4 @@
|
||||
text/
|
||||
trained_models/
|
||||
trec_eval-8.0/trec_eval.dSYM
|
||||
data/
|
||||
|
||||
@@ -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
|
||||
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.
|
||||
@@ -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)
|
||||
|
||||
Executable
+17
@@ -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
|
||||
@@ -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')
|
||||
|
||||
+12
-5
@@ -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)
|
||||
|
||||
@@ -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')
|
||||
@@ -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],
|
||||
|
||||
@@ -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)]
|
||||
)
|
||||
Reference in New Issue
Block a user