From 1b817d3e241dbc2630784e24ec4ddac88c5c3024 Mon Sep 17 00:00:00 2001 From: Achyudh Ram Date: Tue, 2 Oct 2018 21:09:51 -0400 Subject: [PATCH] WIP: Add Reuters-21578 dataset (#147) * Add ReutersTrainer, ReutersEvaluator options in Factory classes * Add Reuters to Kim-CNN command line arguments * Fix SST dataset path according to changes in Kim-CNN args The dataset path in args.py was made to point at the dataset folder rather than dataset/SST folder. Hence SST folder was added to paths in the SST dataset class * Add Reuters dataset class, and support in __main__ * Add Reuters dataset trainers and evaluators * Remove debug print statement in reuters_evaluator * Fix rounding bug in reuters_trainer and reuters_evaluator --- common/evaluation.py | 4 +- common/evaluators/reuters_evaluator.py | 27 +++++++++ common/train.py | 4 +- common/trainers/reuters_trainer.py | 83 ++++++++++++++++++++++++++ datasets/reuters.py | 78 ++++++++++++++++++++++++ datasets/sst.py | 4 +- kim_cnn/__main__.py | 11 ++++ kim_cnn/args.py | 4 +- 8 files changed, 210 insertions(+), 5 deletions(-) create mode 100644 common/evaluators/reuters_evaluator.py create mode 100644 common/trainers/reuters_trainer.py create mode 100644 datasets/reuters.py diff --git a/common/evaluation.py b/common/evaluation.py index 8c7f650..44064ee 100644 --- a/common/evaluation.py +++ b/common/evaluation.py @@ -4,6 +4,7 @@ from .evaluators.sst_evaluator import SSTEvaluator from .evaluators.trecqa_evaluator import TRECQAEvaluator from .evaluators.wikiqa_evaluator import WikiQAEvaluator from .evaluators.pit2015_evaluator import PIT2015Evaluator +from .evaluators.reuters_evaluator import ReutersEvaluator from nce.nce_pairwise_mp.evaluators.trecqa_evaluator import TRECQAEvaluatorNCE from nce.nce_pairwise_mp.evaluators.wikiqa_evaluator import WikiQAEvaluatorNCE @@ -20,7 +21,8 @@ class EvaluatorFactory(object): 'trecqa': TRECQAEvaluator, 'wikiqa': WikiQAEvaluator, 'pit2015': PIT2015Evaluator, - 'twitterurl': PIT2015Evaluator + 'twitterurl': PIT2015Evaluator, + 'Reuters': ReutersEvaluator } evaluator_map_nce = { diff --git a/common/evaluators/reuters_evaluator.py b/common/evaluators/reuters_evaluator.py new file mode 100644 index 0000000..4f84b36 --- /dev/null +++ b/common/evaluators/reuters_evaluator.py @@ -0,0 +1,27 @@ +import torch +import torch.nn.functional as F +import numpy as np + +from .evaluator import Evaluator + + +class ReutersEvaluator(Evaluator): + + def get_scores(self): + self.model.eval() + self.data_loader.init_epoch() + n_dev_correct = 0 + total_loss = 0 + + for batch_idx, batch in enumerate(self.data_loader): + scores = self.model(batch.text) + # Using binary accuracy + for tensor1, tensor2 in zip(F.sigmoid(scores).round().long(), batch.label): + if np.array_equal(tensor1, tensor2): + n_dev_correct += 1 + total_loss += F.binary_cross_entropy_with_logits(scores, batch.label.float(), size_average=False).item() + + accuracy = 100. * n_dev_correct / len(self.data_loader.dataset.examples) + avg_loss = total_loss / len(self.data_loader.dataset.examples) + + return [accuracy, avg_loss], ['accuracy', 'cross_entropy_loss'] diff --git a/common/train.py b/common/train.py index 3099dfd..bfd29df 100644 --- a/common/train.py +++ b/common/train.py @@ -4,6 +4,7 @@ from .trainers.trecqa_trainer import TRECQATrainer from .trainers.wikiqa_trainer import WikiQATrainer from .trainers.pit2015_trainer import PIT2015Trainer from .trainers.sst_trainer import SSTTrainer +from .trainers.reuters_trainer import ReutersTrainer from nce.nce_pairwise_mp.trainers.trecqa_trainer import TRECQATrainerNCE from nce.nce_pairwise_mp.trainers.wikiqa_trainer import WikiQATrainerNCE @@ -20,7 +21,8 @@ class TrainerFactory(object): 'trecqa': TRECQATrainer, 'wikiqa': WikiQATrainer, 'pit2015': PIT2015Trainer, - 'twitterurl': PIT2015Trainer + 'twitterurl': PIT2015Trainer, + 'Reuters': ReutersTrainer } trainer_map_nce = { diff --git a/common/trainers/reuters_trainer.py b/common/trainers/reuters_trainer.py new file mode 100644 index 0000000..1d310b2 --- /dev/null +++ b/common/trainers/reuters_trainer.py @@ -0,0 +1,83 @@ +import time +import os + +import torch +import torch.nn.functional as F +import numpy as np + +from .trainer import Trainer +from utils.serialization import save_checkpoint + + +class ReutersTrainer(Trainer): + + def __init__(self, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator): + super(ReutersTrainer, self).__init__(model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + self.early_stop = False + self.best_dev_acc = 0 + self.iterations = 0 + self.iters_not_improved = 0 + self.start = None + self.log_template = ' '.join( + '{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{},{:12.4f},{}'.split(',')) + self.dev_log_template = ' '.join('{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{:8.6f},{:12.4f},{:12.4f}'.split(',')) + + def train_epoch(self, epoch): + self.train_loader.init_epoch() + n_correct, n_total = 0, 0 + for batch_idx, batch in enumerate(self.train_loader): + self.iterations += 1 + self.model.train() + self.optimizer.zero_grad() + scores = self.model(batch.text) + # Using binary accuracy + for tensor1, tensor2 in zip(F.sigmoid(scores).round().long(), batch.label): + if np.array_equal(tensor1, tensor2): + n_correct += 1 + n_total += batch.batch_size + train_acc = 100. * n_correct / n_total + loss = F.binary_cross_entropy_with_logits(scores, batch.label.float()) + loss.backward() + + self.optimizer.step() + + # Evaluate performance on validation set + if self.iterations % self.dev_log_interval == 1: + dev_acc, dev_loss = self.dev_evaluator.get_scores()[0] + print(self.dev_log_template.format(time.time() - self.start, + epoch, self.iterations, 1 + batch_idx, len(self.train_loader), + 100. * (1 + batch_idx) / len(self.train_loader), loss.item(), + dev_loss, train_acc, dev_acc)) + + # Update validation results + if dev_acc > self.best_dev_acc: + self.iters_not_improved = 0 + self.best_dev_acc = dev_acc + snapshot_path = os.path.join(self.model_outfile, self.train_loader.dataset.NAME, self.model.mode + '_best_model.pt') + torch.save(self.model, snapshot_path) + else: + self.iters_not_improved += 1 + if self.iters_not_improved >= self.patience: + self.early_stop = True + break + + if self.iterations % self.log_interval == 1: + # print progress message + print(self.log_template.format(time.time() - self.start, + epoch, self.iterations, 1 + batch_idx, len(self.train_loader), + 100. * (1 + batch_idx) / len(self.train_loader), loss.item(), ' ' * 8, + train_acc, ' ' * 12)) + + def train(self, epochs): + self.start = time.time() + header = ' Time Epoch Iteration Progress (%Epoch) Loss Dev/Loss Accuracy Dev/Accuracy' + # model_outfile is actually a directory, using model_outfile to conform to Trainer naming convention + os.makedirs(self.model_outfile, exist_ok=True) + os.makedirs(os.path.join(self.model_outfile, self.train_loader.dataset.NAME), exist_ok=True) + print(header) + + for epoch in range(1, epochs + 1): + if self.early_stop: + print("Early Stopping. Epoch: {}, Best Dev Acc: {}".format(epoch, self.best_dev_acc)) + break + self.train_epoch(epoch) diff --git a/datasets/reuters.py b/datasets/reuters.py new file mode 100644 index 0000000..04033b2 --- /dev/null +++ b/datasets/reuters.py @@ -0,0 +1,78 @@ +import re +import os + +import torch +from torchtext.data import Field, TabularDataset +from torchtext.data.iterator import BucketIterator +from torchtext.vocab import Vectors + + +def clean_string(string): + """ + Performs tokenization and string cleaning for the Reuters dataset + """ + string = re.sub(r"[^A-Za-z0-9(),!?\'`]", " ", string) + string = re.sub(r"\s{2,}", " ", string) + return string.lower().strip().split() + + +def clean_string_fl(string): + """ + Returns only the title and first line (excluding the title) for every Reuters article, then calls clean_string + """ + split_string = string.split('.') + if len(split_string) > 1: + return clean_string(split_string[0] + ". " + split_string[1]) + else: + return clean_string(string) + + +def process_labels(string): + """ + Returns the label string as a list of integers + :param string: + :return: + """ + return [float(x) for x in string] + + +class Reuters(TabularDataset): + NAME = 'Reuters' + NUM_CLASSES = 90 + + TEXT_FIELD = Field(batch_first=True, tokenize=clean_string_fl) + LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True, preprocessing=process_labels) + + @staticmethod + def sort_key(ex): + return len(ex.text) + + @classmethod + def splits(cls, path, train=os.path.join('Reuters-21578', 'data', 'reuters_train.tsv'), + validation=os.path.join('Reuters-21578', 'data', 'reuters_validation.tsv'), + test=os.path.join('Reuters-21578', 'data','reuters_test.tsv'), **kwargs): + return super(Reuters, cls).splits( + path, train=train, validation=validation, test=test, + format='tsv', fields=[('label', cls.LABEL_FIELD), ('text', cls.TEXT_FIELD)] + ) + + @classmethod + def iters(cls, path, vectors_name, vectors_cache, batch_size=64, shuffle=True, device=0, vectors=None, + unk_init=torch.Tensor.zero_): + """ + :param path: directory containing train, test, dev files + :param vectors_name: name of word vectors file + :param vectors_cache: path to 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: + """ + if vectors is None: + vectors = Vectors(name=vectors_name, cache=vectors_cache, unk_init=unk_init) + + train, val, test = cls.splits(path) + cls.TEXT_FIELD.build_vocab(train, val, test, vectors=vectors) + return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle, + sort_within_batch=True, device=device) \ No newline at end of file diff --git a/datasets/sst.py b/datasets/sst.py index 6a93f8b..f1c9ed2 100644 --- a/datasets/sst.py +++ b/datasets/sst.py @@ -1,3 +1,4 @@ +import os import re import torch @@ -68,7 +69,8 @@ class SST2(TabularDataset): return len(ex.text) @classmethod - def splits(cls, path, train='stsa.binary.phrases.train', validation='stsa.binary.dev', test='stsa.binary.test', **kwargs): + def splits(cls, path, train=os.path.join('SST', 'stsa.binary.phrases.train'), + validation=os.path.join('SST', 'stsa.binary.dev'), test=os.path.join('SST', 'stsa.binary.test'), **kwargs): return super(SST2, cls).splits( path, train=train, validation=validation, test=test, format='tsv', fields=[('label', cls.LABEL_FIELD), ('text', cls.TEXT_FIELD)] diff --git a/kim_cnn/__main__.py b/kim_cnn/__main__.py index d1c76f0..9c871c0 100644 --- a/kim_cnn/__main__.py +++ b/kim_cnn/__main__.py @@ -10,9 +10,11 @@ from common.evaluation import EvaluatorFactory from common.train import TrainerFactory from datasets.sst import SST1 from datasets.sst import SST2 +from datasets.reuters import Reuters from kim_cnn.args import get_args from kim_cnn.model import KimCNN + class UnknownWordVecCache(object): """ Caches the first randomly generated word vector for a certain size to make it is reused. @@ -76,6 +78,8 @@ if __name__ == '__main__': # Set up the data for training SST-2 elif args.dataset == 'SST-2': train_iter, dev_iter, test_iter = SST2.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk) + elif args.dataset == 'Reuters': + train_iter, dev_iter, test_iter = Reuters.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk) else: raise ValueError('Unrecognized dataset') @@ -113,6 +117,10 @@ if __name__ == '__main__': train_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, train_iter, args.batch_size, args.gpu) test_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, test_iter, args.batch_size, args.gpu) dev_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, dev_iter, args.batch_size, args.gpu) + elif args.dataset == 'Reuters': + train_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, train_iter, args.batch_size, args.gpu) + test_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, test_iter, args.batch_size, args.gpu) + dev_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, dev_iter, args.batch_size, args.gpu) else: raise ValueError('Unrecognized dataset') @@ -141,6 +149,9 @@ if __name__ == '__main__': elif args.dataset == 'SST-2': evaluate_dataset('dev', SST2, model, None, dev_iter, args.batch_size, args.gpu) evaluate_dataset('test', SST2, model, None, test_iter, args.batch_size, args.gpu) + elif args.dataset == 'Reuters': + evaluate_dataset('dev', Reuters, model, None, dev_iter, args.batch_size, args.gpu) + evaluate_dataset('test', Reuters, model, None, test_iter, args.batch_size, args.gpu) else: raise ValueError('Unrecognized dataset') diff --git a/kim_cnn/args.py b/kim_cnn/args.py index 1c42b6e..5eeb55d 100644 --- a/kim_cnn/args.py +++ b/kim_cnn/args.py @@ -12,7 +12,7 @@ def get_args(): parser.add_argument('--mode', type=str, default='multichannel', choices=['rand', 'static', 'non-static', 'multichannel']) parser.add_argument('--lr', type=float, default=1.0) parser.add_argument('--seed', type=int, default=3435) - parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2']) + parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2', 'Reuters']) 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) @@ -24,7 +24,7 @@ def get_args(): parser.add_argument('--dropout', type=float, default=0.5) parser.add_argument('--epoch_decay', type=int, default=15) parser.add_argument('--data_dir', help='word vectors directory', - default=os.path.join(os.pardir, 'Castor-data', 'datasets', 'SST')) + default=os.path.join(os.pardir, 'Castor-data', 'datasets')) parser.add_argument('--word_vectors_dir', help='word vectors directory', default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'word2vec')) parser.add_argument('--word_vectors_file', help='word vectors filename', default='GoogleNews-vectors-negative300.txt')