diff --git a/sm-model/README.md b/sm-model/README.md index 78df4aa..9e1ba8b 100644 --- a/sm-model/README.md +++ b/sm-model/README.md @@ -20,7 +20,7 @@ $ make ``2.`` Get the Overlapping features for Q and A: ``` -$ python overlap_features.py TrecQA +$ python overlap_features.py ../../data/TrecQA ``` ``3.`` To run the S&M model on TrecQA, please follow the same parameter setting: diff --git a/sm-model/main.py b/sm-model/main.py index 758cb47..4286a41 100644 --- a/sm-model/main.py +++ b/sm-model/main.py @@ -67,14 +67,28 @@ if __name__ == "__main__": ap.add_argument('word_vectors_file', help='NOTE: a cache will be created for faster loading for word vectors') ap.add_argument('dataset_folder', help='directory containing train, dev, test sets') ap.add_argument('model_fname', help='model will be saved in args.dataset_folder/') + ap.add_argument('--classes', type=int, default=2) + + # system arguments + # TODO: add arguments for CUDA + ap.add_argument('--num_threads', help="the number of simultaneous processes to run", type=int, default=4) + + # training arguments ap.add_argument('--batch_size', type=int, default=1) - ap.add_argument('--filter_width', type=int, default=5) - ap.add_argument('--epochs', type=int, default=25) + ap.add_argument('--filter_width', type=int, default=5) ap.add_argument('--eta', help='Initial learning rate', default=0.01, type=float) ap.add_argument('--mom', help='SGD Momentum', default=0.9, type=float) - ap.add_argument('--classes', type=int, default=2) + + # epoch related arguments + ap.add_argument('--epochs', type=int, default=25) ap.add_argument('--patience', type=int, default=5, help="if there is no appreciable change in model after epochs, then stop") + # debugging arguments + ap.add_argument('--debugSingleBatch', action="store_true", help="will stop program after training 1 input batch") + ap.add_argument('--no_ext_feats', action="store_true", help="will not include external features in the model") + ap.add_argument('--num_conv_filters', help="the number of convolution channels (lesser is faster)", default=100, type=int) + + args = ap.parse_args() torch.manual_seed(1234) @@ -87,18 +101,20 @@ if __name__ == "__main__": vocab_size, vec_dim = utils.load_embedding_dimensions(cache_file) # instantiate model - net = QAModel(vec_dim, args.filter_width) #filter width is 5 + net = QAModel(vec_dim, args.filter_width, args.num_conv_filters, args.no_ext_feats) #filter width is 5 QAModel.save(net, args.dataset_folder, args.model_fname) + + torch.set_num_threads(args.num_threads) - trainer = Trainer(net) + trainer = Trainer(net, args.eta, args.mom) best_accuracy = 0.0 best_model = 0 for i in range(args.epochs): logger.info('Training epoch {} -------------'.format(i+1)) - train_accuracy = trainer.train(args.dataset_folder, 'train', args.batch_size, cache_file) - # sys.exit(0) + train_accuracy = trainer.train(args.dataset_folder, 'train', args.batch_size, cache_file, args.debugSingleBatch) + if args.debugSingleBatch: sys.exit(0) dev_accuracy, dev_scores = trainer.test(args.dataset_folder, 'clean-dev', args.batch_size, cache_file) if dev_accuracy > best_accuracy: best_model = i diff --git a/sm-model/model.py b/sm-model/model.py index 4bd7f29..5b301e2 100644 --- a/sm-model/model.py +++ b/sm-model/model.py @@ -12,7 +12,7 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) ch = logging.StreamHandler() -ch.setLevel(logging.INFO) +ch.setLevel(logging.DEBUG) formatter = logging.Formatter('%(levelname)s - %(message)s') ch.setFormatter(formatter) logger.addHandler(ch) @@ -28,10 +28,12 @@ class QAModel(nn.Module): def load(in_folder, model_fname): return torch.load(os.path.join(in_folder, model_fname)) - def __init__(self, input_n_dim, filter_width, ext_feats_size=4, n_classes=2): + def __init__(self, input_n_dim, filter_width, conv_filters=100, no_ext_feats=False, ext_feats_size=4, n_classes=2): super(QAModel, self).__init__() - self.conv_channels = 100 + self.no_ext_feats = no_ext_feats + + self.conv_channels = conv_filters n_hidden = 2*self.conv_channels + 1 self.conv_q = nn.Sequential( @@ -44,7 +46,7 @@ class QAModel(nn.Module): nn.Tanh() ) - self.combined_feature_vector = nn.Linear(2*self.conv_channels+ext_feats_size, n_hidden) + self.combined_feature_vector = nn.Linear(2*self.conv_channels + (0 if no_ext_feats else ext_feats_size), n_hidden) #TODO: add +1 to Linear layer^. Will need change in forward function self.combined_features_activation = nn.Tanh() self.dropout = nn.Dropout(0.5) @@ -63,8 +65,15 @@ class QAModel(nn.Module): a = F.max_pool1d(a, a.size()[2]) a = a.view(-1, self.conv_channels) - x = torch.cat([q, a, ext_feats], 1) - # logger.debug('featvec x: {}'.format(x)) + x = None + if self.no_ext_feats: + x = torch.cat([q, a], 1) + logger.debug('no_ext_feats') + else: + x = torch.cat([q, a, ext_feats], 1) + logger.debug('with ext_feats') + + logger.debug('featvec x: {}'.format(x)) # logger.debug(x.creator) x = self.combined_feature_vector.forward(x) @@ -73,9 +82,6 @@ class QAModel(nn.Module): x = self.hidden(x) x = self.logsoftmax(x) - logger.debug('x data {}'.format(x.data)) - logger.debug('x grad {}'.format(x.grad)) - return x diff --git a/sm-model/train.py b/sm-model/train.py index 7f4edad..74b5eaa 100644 --- a/sm-model/train.py +++ b/sm-model/train.py @@ -27,11 +27,11 @@ logger.addHandler(ch) class Trainer(object): - def __init__(self, model): + def __init__(self, model, eta, mom): self.reg = 1e-5 self.model = model self.criterion = nn.CrossEntropyLoss() - self.optimizer = optim.SGD(self.model.parameters(), lr=0.001, weight_decay=self.reg) + self.optimizer = optim.SGD(self.model.parameters(), lr=eta, momentum=mom, weight_decay=self.reg) def regularize_loss(self, loss): @@ -150,7 +150,7 @@ class Trainer(object): return float(total_correct)/len(labels), y_pred - def train(self, dataset_folder, set_folder, batch_size, word_vectors_cache_file): + def train(self, dataset_folder, set_folder, batch_size, word_vectors_cache_file, debugSingleBatch): # read in training data questions, sentences, labels, vocab, maxlen_q, maxlen_s, ext_feats = \ utils.read_in_dataset(dataset_folder, set_folder) @@ -186,7 +186,7 @@ class Trainer(object): # logger.debug('batch_loss {}, batch_correct {}'.format(batch_loss, batch_correct)) train_loss += batch_loss train_correct += batch_correct - # break + if debugSingleBatch: break logger.info('train_correct {}'.format(train_correct)) logger.info('train_loss {}'.format(train_loss))