From ab7b667ce2335f7bb1d2a2ea75e821334d05f88b Mon Sep 17 00:00:00 2001 From: Michael Tu Date: Sat, 3 Feb 2018 19:39:08 -0500 Subject: [PATCH] Add twitter dataset --- datasets/twitter.py | 99 +++++++++++++++++++++++++++++++++++++++++++++ mp_cnn/dataset.py | 11 ++++- mp_cnn/main.py | 18 +++++++-- 3 files changed, 124 insertions(+), 4 deletions(-) create mode 100644 datasets/twitter.py diff --git a/datasets/twitter.py b/datasets/twitter.py new file mode 100644 index 0000000..b80d316 --- /dev/null +++ b/datasets/twitter.py @@ -0,0 +1,99 @@ +import os + +import torch + +from torchtext.data.dataset import Dataset +from torchtext.data.example import Example +from torchtext.data.field import Field +from torchtext.data.iterator import BucketIterator +from torchtext.data.iterator import Iterator +from torchtext.vocab import Vectors +from torchtext.data import Pipeline + +from datasets.castor_dataset import CastorPairDataset +from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features + + +class TWITTER(Dataset): + NAME = 'twitter' + NUM_CLASSES = 2 + ID_FIELD = Field(sequential=False, tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True) + AID_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + TEXT_FIELD = Field(batch_first=True, tokenize=lambda x: x) # tokenizer is identity since we already tokenized it to compute external features + EXT_FEATS_FIELD = Field(tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True, tokenize=lambda x: x, + postprocessing=Pipeline(lambda arr, _, train: [float(y) for y in arr])) + LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + VOCAB_SIZE = 0 + + @staticmethod + def sort_key(ex): + return len(ex.sentence_1) + + def __init__(self, pardir, subdirs): + """ + Create a Twitter dataset instance. + """ + fields = [('id', self.ID_FIELD), ('sentence_1', self.TEXT_FIELD), ('sentence_2', self.TEXT_FIELD), ('ext_feats', + self.EXT_FEATS_FIELD), ('label', self.LABEL_FIELD), ('aid', self.AID_FIELD)] + + examples = [] + for subdir in subdirs: + path = os.path.join(pardir, subdir) + with open(os.path.join(path, 'a.toks'), 'r') as f1, open(os.path.join(path, 'b.toks'), 'r') as f2: + sent_list_1 = [l.rstrip('.\n').split(' ') for l in f1] + sent_list_2 = [l.rstrip('.\n').split(' ') for l in f2] + + word_to_doc_cnt = get_pairwise_word_to_doc_freq(sent_list_1, sent_list_2) + overlap_feats = get_pairwise_overlap_features(sent_list_1, sent_list_2, word_to_doc_cnt) + + for subdir_i, subdir in enumerate(subdirs): + path = os.path.join(pardir, subdir) + with open(os.path.join(path, 'id.txt'), 'r') as id_file, open(os.path.join(path, 'sim.txt'), 'r') as label_file: + for i, (pair_id, l1, l2, ext_feats, label) in enumerate(zip(id_file, sent_list_1, sent_list_2, overlap_feats, label_file)): + pair_id = pair_id.rstrip('.\n') + label = label.rstrip('.\n') + example_list = [pair_id, l1, l2, ext_feats, label, (subdir_i) * 100000 + (i + 1)] + example = Example.fromlist(example_list, fields) + examples.append(example) + + super(TWITTER, self).__init__(examples, fields) + + @classmethod + def splits(cls, path, train_paths, test_paths, **kwargs): + train_data = cls(path, train_paths, **kwargs) + test_data = cls(path, test_paths, **kwargs) + return train_data, test_data + + @classmethod + def set_vectors(cls, field, vector_path): + return CastorPairDataset.set_vectors(field, vector_path) + + @classmethod + def iters(cls, path, train_dirs, test_dirs, vectors_name, vectors_dir, batch_size=64, shuffle=True, device=0, pt_file=False, vectors=None, unk_init=torch.Tensor.zero_): + """ + :param path: directory containing train, test, dev files + :param train_dirs: list of directory names used for training + :param test_dirs: list of directory name used for testing + :param vectors_name: name of word vectors file + :param vectors_dir: directory containing word vectors file + :param batch_size: batch size + :param device: GPU device + :param vectors: custom vectors - either predefined torchtext vectors or your own custom Vector classes + :param unk_init: function used to generate vector for OOV words + :return: + """ + + train, test = cls.splits(path, train_dirs, test_dirs) + if not pt_file: + if vectors is None: + vectors = Vectors(name=vectors_name, cache=vectors_dir, unk_init=unk_init) + cls.TEXT_FIELD.build_vocab(train, test, vectors=vectors) + else: + cls.TEXT_FIELD.build_vocab(train, test) + cls.TEXT_FIELD = cls.set_vectors(cls.TEXT_FIELD, os.path.join(vectors_dir, vectors_name)) + + cls.LABEL_FIELD.build_vocab(train, test) + + cls.VOCAB_SIZE = len(cls.TEXT_FIELD.vocab) + + return BucketIterator.splits((train, test), batch_size=batch_size, repeat=False, shuffle=shuffle, device=device) diff --git a/mp_cnn/dataset.py b/mp_cnn/dataset.py index 0ffb2a0..2252452 100644 --- a/mp_cnn/dataset.py +++ b/mp_cnn/dataset.py @@ -6,6 +6,7 @@ import torch.nn as nn from datasets.sick import SICK from datasets.msrvid import MSRVID from datasets.trecqa import TRECQA +from datasets.twitter import TWITTER from datasets.wikiqa import WikiQA @@ -29,7 +30,7 @@ class MPCNNDatasetFactory(object): Get the corresponding Dataset class for a particular dataset. """ @staticmethod - def get_dataset(dataset_name, word_vectors_dir, word_vectors_file, batch_size, device, castor_dir="../", utils_trecqa="utils/trec_eval-9.0.5/trec_eval"): + def get_dataset(dataset_name, word_vectors_dir, word_vectors_file, batch_size, device, castor_dir="../", utils_trecqa="utils/trec_eval-9.0.5/trec_eval", train_dirs=None, test_dirs=None): if dataset_name == 'sick': dataset_root = os.path.join(os.pardir, castor_dir, 'data', 'sick/') train_loader, dev_loader, test_loader = SICK.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) @@ -63,6 +64,14 @@ class MPCNNDatasetFactory(object): embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) embedding.weight = nn.Parameter(WikiQA.TEXT_FIELD.vocab.vectors) return WikiQA, embedding, train_loader, test_loader, dev_loader + elif dataset_name == 'twitter': + dataset_root = os.path.join(os.pardir, castor_dir, 'data', 'twitter-microblog/order_by_rel/') + dev_loader = None + train_loader, test_loader = TWITTER.iters(dataset_root, train_dirs, test_dirs, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) + embedding_dim = TWITTER.TEXT_FIELD.vocab.vectors.size() + embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) + embedding.weight = nn.Parameter(TWITTER.TEXT_FIELD.vocab.vectors) + return TWITTER, embedding, train_loader, test_loader, dev_loader else: raise ValueError('{} is not a valid dataset.'.format(dataset_name)) diff --git a/mp_cnn/main.py b/mp_cnn/main.py index a91a300..8ba1458 100644 --- a/mp_cnn/main.py +++ b/mp_cnn/main.py @@ -5,6 +5,7 @@ import pprint import random import numpy as np +import sys import torch import torch.optim as optim @@ -17,9 +18,11 @@ from mp_cnn.train import MPCNNTrainerFactory if __name__ == '__main__': parser = argparse.ArgumentParser(description='PyTorch implementation of Multi-Perspective CNN') parser.add_argument('model_outfile', help='file to save final model') - parser.add_argument('--dataset', help='dataset to use, one of [sick, msrvid, trecqa, wikiqa]', default='sick') + parser.add_argument('--dataset', help='dataset to use, one of [sick, msrvid, trecqa, wikiqa, twitter]', default='sick') parser.add_argument('--word-vectors-dir', help='word vectors directory', default=os.path.join(os.pardir, os.pardir, 'data', 'GloVe')) parser.add_argument('--word-vectors-file', help='word vectors filename', default='glove.840B.300d.txt') + parser.add_argument('--train_dirs', nargs='+', help='training directory names for twitter dataset') + parser.add_argument('--test_dirs', nargs='+', help='testing directory names for twitter dataset') parser.add_argument('--skip-training', help='will load pre-trained model', action='store_true') parser.add_argument('--device', type=int, default=0, help='GPU device, -1 for CPU (default: 0)') parser.add_argument('--sparse-features', action='store_true', default=False, help='use sparse features (default: false)') @@ -61,8 +64,17 @@ if __name__ == '__main__': logger.info(pprint.pformat(vars(args))) - dataset_cls, embedding, train_loader, test_loader, dev_loader \ - = MPCNNDatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device) + if args.dataset == 'twitter': + if not args.train_dirs or not args.test_dirs: + print('For twitter dataset --train_dirs and --test_dirs must be specified') + sys.exit(1) + dataset_cls, embedding, train_loader, test_loader, dev_loader \ + = MPCNNDatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device, train_dirs=args.train_dirs, test_dirs=args.test_dirs) + else: + dataset_cls, embedding, train_loader, test_loader, dev_loader \ + = MPCNNDatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device) + + import ipdb; ipdb.set_trace() filter_widths = list(range(1, args.max_window_size + 1)) + [np.inf] model = MPCNN(embedding, args.holistic_filters, args.per_dim_filters, filter_widths,