From 97bdaec5bbbc82ddf967f210d3ebd15ebc038004 Mon Sep 17 00:00:00 2001 From: Ashutosh-Adhikari Date: Fri, 9 Nov 2018 20:53:43 -0500 Subject: [PATCH] Add AAPD for XML_CNN (#160) * Add AAPD for XMLCNN * Add kwargs for XML --- xml_cnn/__main__.py | 52 +++++++++++++++++---------------------------- xml_cnn/args.py | 6 +++--- xml_cnn/model.py | 2 +- 3 files changed, 23 insertions(+), 37 deletions(-) diff --git a/xml_cnn/__main__.py b/xml_cnn/__main__.py index f6de650..283c596 100644 --- a/xml_cnn/__main__.py +++ b/xml_cnn/__main__.py @@ -13,6 +13,7 @@ from common.train import TrainerFactory from datasets.sst import SST1 from datasets.sst import SST2 from datasets.reuters import Reuters +from datasets.aapd import AAPD from xml_cnn.args import get_args from xml_cnn.model import XmlCNN @@ -74,16 +75,16 @@ if __name__ == '__main__': random.seed(args.seed) logger = get_logger() - # Set up the data for training SST-1 - if args.dataset == 'SST-1': - train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk) - # 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: + dataset_map = { + 'SST-1': SST1, + 'SST-2': SST2, + 'Reuters': Reuters, + 'AAPD': AAPD + } + if args.dataset not in dataset_map: raise ValueError('Unrecognized dataset') + else: + train_iter, dev_iter, test_iter = dataset_map[args.dataset].iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk) config = deepcopy(args) config.dataset = train_iter.dataset @@ -112,21 +113,12 @@ if __name__ == '__main__': #optimizer = torch.optim.Adadelta(parameter, lr=args.lr, weight_decay=args.weight_decay) optimizer = torch.optim.Adam(parameter, lr = args.lr) - if args.dataset == 'SST-1': - train_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, train_iter, args.batch_size, args.gpu) - test_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, test_iter, args.batch_size, args.gpu) - dev_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, dev_iter, args.batch_size, args.gpu) - elif args.dataset == 'SST-2': - 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: + if args.dataset not in dataset_map: raise ValueError('Unrecognized dataset') - + else: + train_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, train_iter, args.batch_size, args.gpu) + test_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, test_iter, args.batch_size, args.gpu) + dev_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, dev_iter, args.batch_size, args.gpu) trainer_config = { 'optimizer': optimizer, 'batch_size': args.batch_size, @@ -146,17 +138,11 @@ if __name__ == '__main__': else: model = torch.load(args.trained_model, map_location=lambda storage, location: storage) - if args.dataset == 'SST-1': - evaluate_dataset('dev', SST1, model, None, dev_iter, args.batch_size, args.gpu) - evaluate_dataset('test', SST1, model, None, test_iter, args.batch_size, args.gpu) - 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: + if args.dataset not in dataset_map: raise ValueError('Unrecognized dataset') + else: + evaluate_dataset('dev', dataset_map[args.dataset], model, None, dev_iter, args.batch_size, args.gpu) + evaluate_dataset('test', dataset_map[args.dataset], model, None, test_iter, args.batch_size, args.gpu) # Calculate dev and test metrics for data_loader in [dev_iter, test_iter]: diff --git a/xml_cnn/args.py b/xml_cnn/args.py index b948c01..a72ed4a 100644 --- a/xml_cnn/args.py +++ b/xml_cnn/args.py @@ -4,7 +4,7 @@ from argparse import ArgumentParser def get_args(): - parser = ArgumentParser(description="Kim CNN") + parser = ArgumentParser(description="XML CNN") parser.add_argument('--no_cuda', action='store_false', help='do not use cuda', dest='cuda') parser.add_argument('--gpu', type=int, default=0) # Use -1 for CPU parser.add_argument('--epochs', type=int, default=30) @@ -12,12 +12,12 @@ 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', 'Reuters']) + parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2', 'Reuters','AAPD']) 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) parser.add_argument('--patience', type=int, default=50) - parser.add_argument('--save_path', type=str, default='kim_cnn/saves') + parser.add_argument('--save_path', type=str, default='xml_cnn/saves') parser.add_argument('--output_channel', type=int, default=100) parser.add_argument('--words_dim', type=int, default=300) parser.add_argument('--embed_dim', type=int, default=300) diff --git a/xml_cnn/model.py b/xml_cnn/model.py index 599ee62..6b9c3dc 100644 --- a/xml_cnn/model.py +++ b/xml_cnn/model.py @@ -48,7 +48,7 @@ class XmlCNN(nn.Module): - def forward(self, x): + def forward(self, x, **kwargs): if self.mode == 'rand': word_input = self.embed(x) # (batch, sent_len, embed_dim) x = word_input.unsqueeze(1) # (batch, channel_input, sent_len, embed_dim)