mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Add twitter dataset
This commit is contained in:
@@ -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)
|
||||
+10
-1
@@ -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))
|
||||
|
||||
|
||||
+15
-3
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user