diff --git a/common/dataset.py b/common/dataset.py index 32a76cd..318a232 100644 --- a/common/dataset.py +++ b/common/dataset.py @@ -29,39 +29,31 @@ class DatasetFactory(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"): if dataset_name == 'sick': dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-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) - embedding_dim = SICK.TEXT_FIELD.vocab.vectors.size() - embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) - embedding.weight = nn.Parameter(SICK.TEXT_FIELD.vocab.vectors) + embedding = nn.Embedding.from_pretrained(SICK.TEXT_FIELD.vocab.vectors) return SICK, embedding, train_loader, test_loader, dev_loader elif dataset_name == 'msrvid': dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'msrvid/') dev_loader = None train_loader, test_loader = MSRVID.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) - embedding_dim = MSRVID.TEXT_FIELD.vocab.vectors.size() - embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) - embedding.weight = nn.Parameter(MSRVID.TEXT_FIELD.vocab.vectors) + embedding = nn.Embedding.from_pretrained(MSRVID.TEXT_FIELD.vocab.vectors) return MSRVID, embedding, train_loader, test_loader, dev_loader elif dataset_name == 'trecqa': if not os.path.exists(os.path.join(castor_dir, utils_trecqa)): raise FileNotFoundError('TrecQA requires the trec_eval tool to run. Please run get_trec_eval.sh inside Castor/utils (as working directory) before continuing.') dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'TrecQA/') train_loader, dev_loader, test_loader = TRECQA.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) - embedding_dim = TRECQA.TEXT_FIELD.vocab.vectors.size() - embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) - embedding.weight = nn.Parameter(TRECQA.TEXT_FIELD.vocab.vectors) + embedding = nn.Embedding.from_pretrained(TRECQA.TEXT_FIELD.vocab.vectors) return TRECQA, embedding, train_loader, test_loader, dev_loader elif dataset_name == 'wikiqa': if not os.path.exists(os.path.join(castor_dir, utils_trecqa)): raise FileNotFoundError('TrecQA requires the trec_eval tool to run. Please run get_trec_eval.sh inside Castor/utils (as working directory) before continuing.') dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'WikiQA/') train_loader, dev_loader, test_loader = WikiQA.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) - embedding_dim = WikiQA.TEXT_FIELD.vocab.vectors.size() - embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) - embedding.weight = nn.Parameter(WikiQA.TEXT_FIELD.vocab.vectors) + embedding = nn.Embedding.from_pretrained(WikiQA.TEXT_FIELD.vocab.vectors) return WikiQA, embedding, train_loader, test_loader, dev_loader else: raise ValueError('{} is not a valid dataset.'.format(dataset_name)) diff --git a/common/evaluation.py b/common/evaluation.py index aa34dee..d3d80e7 100644 --- a/common/evaluation.py +++ b/common/evaluation.py @@ -22,7 +22,7 @@ class EvaluatorFactory(object): } @staticmethod - def get_evaluator(dataset_cls, model, data_loader, batch_size, device, nce=False): + def get_evaluator(dataset_cls, model, embedding, data_loader, batch_size, device, nce=False): if data_loader is None: return None @@ -38,5 +38,5 @@ class EvaluatorFactory(object): raise ValueError('{} is not implemented.'.format(dataset_cls)) return evaluator_map[dataset_cls.NAME]( - dataset_cls, model, data_loader, batch_size, device + dataset_cls, model, embedding, data_loader, batch_size, device ) diff --git a/common/evaluators/evaluator.py b/common/evaluators/evaluator.py index ad3fbba..7318bec 100644 --- a/common/evaluators/evaluator.py +++ b/common/evaluators/evaluator.py @@ -1,15 +1,21 @@ class Evaluator(object): """ - Evaluates performance of model on a Dataset, using metrics specific to the Dataset. + Evaluates a model on a Dataset, using metrics specific to the Dataset. """ - def __init__(self, dataset_cls, model, data_loader, batch_size, device): + def __init__(self, dataset_cls, model, embedding, data_loader, batch_size, device): self.dataset_cls = dataset_cls self.model = model + self.embedding = embedding self.data_loader = data_loader self.batch_size = batch_size self.device = device + def get_sentence_embeddings(self, batch): + sent1 = self.embedding(batch.sentence_1).transpose(1, 2) + sent2 = self.embedding(batch.sentence_2).transpose(1, 2) + return sent1, sent2 + def get_scores(self): """ Get the scores used to evaluate the model. diff --git a/common/evaluators/msrvid_evaluator.py b/common/evaluators/msrvid_evaluator.py index 975638e..ebcb08c 100644 --- a/common/evaluators/msrvid_evaluator.py +++ b/common/evaluators/msrvid_evaluator.py @@ -7,30 +7,27 @@ from .evaluator import Evaluator class MSRVIDEvaluator(Evaluator): - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(MSRVIDEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) - def get_scores(self): self.model.eval() num_classes = self.dataset_cls.NUM_CLASSES - predict_classes = torch.arange(0, num_classes).expand(self.batch_size, num_classes) test_kl_div_loss = 0 predictions = [] true_labels = [] for batch in self.data_loader: - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - test_kl_div_loss += F.kl_div(output, batch.label, size_average=False).data[0] + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats) + test_kl_div_loss += F.kl_div(output, batch.label, size_average=False).item() + + predict_classes = batch.label.new_tensor(torch.arange(0, num_classes)).expand(self.batch_size, num_classes) # handle last batch which might have smaller size if len(predict_classes) != len(batch.sentence_1): - predict_classes = torch.arange(0, num_classes).expand(len(batch.sentence_1), num_classes) + predict_classes = batch.label.new_tensor(torch.arange(0, num_classes)).expand(len(batch.sentence_1), num_classes) - if self.data_loader.device != -1: - with torch.cuda.device(self.device): - predict_classes = predict_classes.cuda() - - true_labels.append((predict_classes * batch.label.data).sum(dim=1)) - predictions.append((predict_classes * output.data.exp()).sum(dim=1)) + true_labels.append((predict_classes * batch.label.detach()).sum(dim=1)) + predictions.append((predict_classes * output.detach().exp()).sum(dim=1)) del output @@ -40,3 +37,12 @@ class MSRVIDEvaluator(Evaluator): pearson_r = pearsonr(predictions, true_labels)[0] return [pearson_r, test_kl_div_loss], ['pearson_r', 'KL-divergence loss'] + + def get_final_prediction_and_label(self, batch_predictions, batch_labels): + num_classes = self.dataset_cls.NUM_CLASSES + predict_classes = batch_labels.new_tensor(torch.arange(0, num_classes)).expand(batch_predictions.size(0), num_classes) + + predictions = (predict_classes * batch_predictions.exp()).sum(dim=1) + true_labels = (predict_classes * batch_labels).sum(dim=1) + + return predictions, true_labels diff --git a/common/evaluators/qa_evaluator.py b/common/evaluators/qa_evaluator.py index dc61297..e48b118 100644 --- a/common/evaluators/qa_evaluator.py +++ b/common/evaluators/qa_evaluator.py @@ -6,9 +6,6 @@ from utils.relevancy_metrics import get_map_mrr class QAEvaluator(Evaluator): - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(QAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) - def get_scores(self): self.model.eval() test_cross_entropy_loss = 0 @@ -17,12 +14,15 @@ class QAEvaluator(Evaluator): predictions = [] for batch in self.data_loader: - qids.extend(batch.id.data.cpu().numpy()) - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - test_cross_entropy_loss += F.cross_entropy(output, batch.label, size_average=False).data[0] + qids.extend(batch.id.detach().cpu().numpy()) + # Select embedding + sent1, sent2, sent1_nonstatic, sent2_nonstatic = self.get_sentence_embeddings(batch) - true_labels.extend(batch.label.data.cpu().numpy()) - predictions.extend(output.data.exp()[:, 1].cpu().numpy()) + output = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw, sent1_nonstatic, sent2_nonstatic) + test_cross_entropy_loss += F.cross_entropy(output, batch.label, size_average=False).item() + + true_labels.extend(batch.label.detach().cpu().numpy()) + predictions.extend(output.detach().exp()[:, 1].cpu().numpy()) del output @@ -31,4 +31,9 @@ class QAEvaluator(Evaluator): mean_average_precision, mean_reciprocal_rank = get_map_mrr(qids, predictions, true_labels, self.data_loader.device) test_cross_entropy_loss /= len(batch.dataset.examples) - return [test_cross_entropy_loss, mean_average_precision, mean_reciprocal_rank], ['cross entropy loss', 'map', 'mrr'] + return [mean_average_precision, mean_reciprocal_rank, test_cross_entropy_loss], ['map', 'mrr', 'cross entropy loss'] + + def get_final_prediction_and_label(self, batch_predictions, batch_labels): + predictions = batch_predictions.exp()[:, 1] + + return predictions, batch_labels \ No newline at end of file diff --git a/common/evaluators/sick_evaluator.py b/common/evaluators/sick_evaluator.py index d0f2280..4cf9ce8 100644 --- a/common/evaluators/sick_evaluator.py +++ b/common/evaluators/sick_evaluator.py @@ -7,37 +7,46 @@ from .evaluator import Evaluator class SICKEvaluator(Evaluator): - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(SICKEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) - def get_scores(self): self.model.eval() num_classes = self.dataset_cls.NUM_CLASSES - predict_classes = torch.arange(1, num_classes + 1).expand(self.batch_size, num_classes) test_kl_div_loss = 0 predictions = [] true_labels = [] for batch in self.data_loader: - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - test_kl_div_loss += F.kl_div(output, batch.label, size_average=False).data[0] + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw) + test_kl_div_loss += F.kl_div(output, batch.label, size_average=False).item() + + predict_classes = batch.label.new_tensor(torch.arange(1, num_classes + 1)).expand(self.batch_size, num_classes) # handle last batch which might have smaller size if len(predict_classes) != len(batch.sentence_1): - predict_classes = torch.arange(1, num_classes + 1).expand(len(batch.sentence_1), num_classes) + predict_classes = batch.label.new_tensor(torch.arange(1, num_classes + 1)).expand(len(batch.sentence_1), num_classes) - if self.data_loader.device != -1: - with torch.cuda.device(self.device): - predict_classes = predict_classes.cuda() - - true_labels.append((predict_classes * batch.label.data).sum(dim=1)) - predictions.append((predict_classes * output.data.exp()).sum(dim=1)) + true_labels.append((predict_classes * batch.label.detach()).sum(dim=1)) + predictions.append((predict_classes * output.detach().exp()).sum(dim=1)) del output - predictions = torch.cat(predictions).cpu().numpy() - true_labels = torch.cat(true_labels).cpu().numpy() + predictions = torch.cat(predictions) + true_labels = torch.cat(true_labels) + mse = F.mse_loss(predictions, true_labels).item() test_kl_div_loss /= len(batch.dataset.examples) + predictions = predictions.cpu().numpy() + true_labels = true_labels.cpu().numpy() pearson_r = pearsonr(predictions, true_labels)[0] spearman_r = spearmanr(predictions, true_labels)[0] - return [pearson_r, spearman_r, test_kl_div_loss], ['pearson_r', 'spearman_r', 'KL-divergence loss'] + return [pearson_r, spearman_r, mse, test_kl_div_loss], ['pearson_r', 'spearman_r', 'mse', 'KL-divergence loss'] + + def get_final_prediction_and_label(self, batch_predictions, batch_labels): + num_classes = self.dataset_cls.NUM_CLASSES + predict_classes = batch_labels.new_tensor(torch.arange(1, num_classes + 1)).expand(batch_predictions.size(0), num_classes) + + predictions = (predict_classes * batch_predictions.exp()).sum(dim=1) + true_labels = (predict_classes * batch_labels).sum(dim=1) + + return predictions, true_labels diff --git a/common/evaluators/trecqa_evaluator.py b/common/evaluators/trecqa_evaluator.py index 93b1d4a..04c0e29 100644 --- a/common/evaluators/trecqa_evaluator.py +++ b/common/evaluators/trecqa_evaluator.py @@ -2,6 +2,4 @@ from .qa_evaluator import QAEvaluator class TRECQAEvaluator(QAEvaluator): - - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(TRECQAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) + pass diff --git a/common/evaluators/wikiqa_evaluator.py b/common/evaluators/wikiqa_evaluator.py index 810bacb..a5ff251 100644 --- a/common/evaluators/wikiqa_evaluator.py +++ b/common/evaluators/wikiqa_evaluator.py @@ -2,6 +2,4 @@ from .qa_evaluator import QAEvaluator class WikiQAEvaluator(QAEvaluator): - - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(WikiQAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) + pass diff --git a/common/train.py b/common/train.py index 96810ec..7d01c37 100644 --- a/common/train.py +++ b/common/train.py @@ -23,7 +23,7 @@ class TrainerFactory(object): } @staticmethod - def get_trainer(dataset_name, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None, nce=False): + def get_trainer(dataset_name, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None, nce=False): if nce: trainer_map = TrainerFactory.trainer_map_nce else: @@ -33,5 +33,5 @@ class TrainerFactory(object): raise ValueError('{} is not implemented.'.format(dataset_name)) return trainer_map[dataset_name]( - model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator + model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator ) diff --git a/common/trainers/msrvid_trainer.py b/common/trainers/msrvid_trainer.py index 7390c77..68a6212 100644 --- a/common/trainers/msrvid_trainer.py +++ b/common/trainers/msrvid_trainer.py @@ -7,13 +7,11 @@ from torch.optim.lr_scheduler import ReduceLROnPlateau from scipy.stats import pearsonr from .trainer import Trainer +from utils.serialization import save_checkpoint class MSRVIDTrainer(Trainer): - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): - super(MSRVIDTrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) - def train_epoch(self, epoch): self.model.train() total_loss = 0 @@ -34,22 +32,26 @@ class MSRVIDTrainer(Trainer): left_out_val_labels.append(batch.label) continue self.optimizer.zero_grad() - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - loss = F.kl_div(output, batch.label) - total_loss += loss.data[0] + + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats) + loss = F.kl_div(output, batch.label, size_average=False) + total_loss += loss.item() loss.backward() self.optimizer.step() if batch_idx % self.log_interval == 0: self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( epoch, min(batch_idx * self.batch_size, len(batch.dataset.examples)), len(batch.dataset.examples), - 100. * batch_idx / (len(self.train_loader)), loss.data[0]) + 100. * batch_idx / (len(self.train_loader)), loss.item() / len(batch)) ) self.evaluate(self.train_evaluator, 'train') if self.use_tensorboard: - self.writer.add_scalar('msrvid/train/kl_div_loss', total_loss, epoch) + self.writer.add_scalar('msrvid/train/kl_div_loss', total_loss / len(self.train_loader.dataset.examples), epoch) return left_out_val_a, left_out_val_b, left_out_val_ext_feats, left_out_val_labels @@ -67,15 +69,17 @@ class MSRVIDTrainer(Trainer): all_predictions, all_true_labels = [], [] val_kl_div_loss = 0 for i in range(len(left_out_a)): - output = self.model(left_out_a[i], left_out_b[i], left_out_ext_feats[i]) - val_kl_div_loss += F.kl_div(output, left_out_label[i], size_average=False).data[0] - predict_classes = torch.arange(0, self.train_loader.dataset.NUM_CLASSES).expand(len(left_out_a[i]), self.train_loader.dataset.NUM_CLASSES) - if self.train_loader.device != -1: - with torch.cuda.device(self.train_loader.device): - predict_classes = predict_classes.cuda() + # Select embedding + sent1 = self.embedding(left_out_a[i]).transpose(1, 2) + sent2 = self.embedding(left_out_b[i]).transpose(1, 2) - predictions = (predict_classes * output.data.exp()).sum(dim=1) - true_labels = (predict_classes * left_out_label[i].data).sum(dim=1) + output = self.model(sent1, sent2, left_out_ext_feats[i]) + val_kl_div_loss += F.kl_div(output, left_out_label[i], size_average=False).item() + predict_classes = left_out_a[i].new_tensor(torch.arange(0, self.train_loader.dataset.NUM_CLASSES))\ + .float().expand(len(left_out_a[i]), self.train_loader.dataset.NUM_CLASSES) + + predictions = (predict_classes * output.detach().exp()).sum(dim=1) + true_labels = (predict_classes * left_out_label[i].detach()).sum(dim=1) all_predictions.append(predictions) all_true_labels.append(true_labels) @@ -88,7 +92,7 @@ class MSRVIDTrainer(Trainer): self.writer.add_scalar('msrvid/dev/pearson_r', pearson_r, epoch) for param_group in self.optimizer.param_groups: - self.logger.info('Validation size: %s Pearson\'s r: %s', output.size()[0], pearson_r) + self.logger.info('Validation size: %s Pearson\'s r: %s', output.size(0), pearson_r) self.logger.info('Learning rate: %s', param_group['lr']) if self.use_tensorboard: @@ -105,7 +109,7 @@ class MSRVIDTrainer(Trainer): if pearson_r > best_dev_score: best_dev_score = pearson_r - torch.save(self.model, self.model_outfile) + save_checkpoint(epoch, self.model.arch, self.model.state_dict(), self.optimizer.state_dict(), best_dev_score, self.model_outfile) if abs(prev_loss - val_kl_div_loss) <= 0.0005: self.logger.info('Early stopping. Loss changed by less than 0.0005.') @@ -114,4 +118,4 @@ class MSRVIDTrainer(Trainer): prev_loss = val_kl_div_loss self.evaluate(self.test_evaluator, 'test') - self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) + self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) \ No newline at end of file diff --git a/common/trainers/qa_trainer.py b/common/trainers/qa_trainer.py index 7960b8c..bc36872 100644 --- a/common/trainers/qa_trainer.py +++ b/common/trainers/qa_trainer.py @@ -1,35 +1,36 @@ import time -import torch import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau from .trainer import Trainer +from utils.serialization import save_checkpoint class QATrainer(Trainer): - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): - super(QATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) - def train_epoch(self, epoch): self.model.train() total_loss = 0 for batch_idx, batch in enumerate(self.train_loader): self.optimizer.zero_grad() - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - loss = F.cross_entropy(output, batch.label, size_average=False) - total_loss += loss.data[0] + + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw) + loss = F.nll_loss(output, batch.label, size_average=False) + total_loss += loss.item() loss.backward() self.optimizer.step() if batch_idx % self.log_interval == 0: self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( epoch, min(batch_idx * self.batch_size, len(batch.dataset.examples)), len(batch.dataset.examples), - 100. * batch_idx / (len(self.train_loader)), loss.data[0]) + 100. * batch_idx / (len(self.train_loader)), loss.item() / len(batch)) ) - average_loss, mean_average_precision, mean_reciprocal_rank = self.evaluate(self.train_evaluator, 'train') + mean_average_precision, mean_reciprocal_rank, average_loss = self.evaluate(self.train_evaluator, 'train') if self.use_tensorboard: self.writer.add_scalar('{}/train/cross_entropy_loss'.format(self.train_loader.dataset.NAME), average_loss, epoch) @@ -49,7 +50,7 @@ class QATrainer(Trainer): self.train_epoch(epoch) dev_scores = self.evaluate(self.dev_evaluator, 'dev') - new_loss, mean_average_precision, mean_reciprocal_rank = dev_scores + mean_average_precision, mean_reciprocal_rank, new_loss = dev_scores if self.use_tensorboard: self.writer.add_scalar('{}/lr'.format(self.train_loader.dataset.NAME), self.optimizer.param_groups[0]['lr'], epoch) @@ -62,15 +63,15 @@ class QATrainer(Trainer): self.logger.info('Epoch {} finished in {:.2f} minutes'.format(epoch, duration / 60)) epoch_times.append(duration) - if dev_scores[0] > best_dev_score: - best_dev_score = dev_scores[0] - torch.save(self.model, self.model_outfile) + if mean_average_precision > best_dev_score: + best_dev_score = mean_average_precision + save_checkpoint(epoch, self.model.arch, self.model.state_dict(), self.optimizer.state_dict(), best_dev_score, self.model_outfile) if abs(prev_loss - new_loss) <= 0.0002: self.logger.info('Early stopping. Loss changed by less than 0.0002.') break prev_loss = new_loss - scheduler.step(dev_scores[0]) + scheduler.step(mean_average_precision) self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) diff --git a/common/trainers/sick_trainer.py b/common/trainers/sick_trainer.py index c0e5d96..ca50e24 100644 --- a/common/trainers/sick_trainer.py +++ b/common/trainers/sick_trainer.py @@ -1,36 +1,37 @@ import time -import torch import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau from .trainer import Trainer +from utils.serialization import save_checkpoint class SICKTrainer(Trainer): - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): - super(SICKTrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) - def train_epoch(self, epoch): self.model.train() total_loss = 0 for batch_idx, batch in enumerate(self.train_loader): self.optimizer.zero_grad() - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - loss = F.kl_div(output, batch.label) - total_loss += loss.data[0] + + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw) + loss = F.kl_div(output, batch.label, size_average=False) + total_loss += loss.item() loss.backward() self.optimizer.step() if batch_idx % self.log_interval == 0: self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( epoch, min(batch_idx * self.batch_size, len(batch.dataset.examples)), len(batch.dataset.examples), - 100. * batch_idx / (len(self.train_loader)), loss.data[0]) + 100. * batch_idx / (len(self.train_loader)), loss.item() / len(batch)) ) if self.use_tensorboard: - self.writer.add_scalar('sick/train/kl_div_loss', total_loss, epoch) + self.writer.add_scalar('sick/train/kl_div_loss', total_loss / len(self.train_loader.dataset.examples), epoch) return total_loss @@ -44,12 +45,11 @@ class SICKTrainer(Trainer): self.logger.info('Epoch {} started...'.format(epoch)) self.train_epoch(epoch) - dev_scores = self.evaluate(self.dev_evaluator, 'dev') - new_loss = dev_scores[2] + pearson, spearman, mse, new_loss = self.evaluate(self.dev_evaluator, 'dev') if self.use_tensorboard: self.writer.add_scalar('sick/lr', self.optimizer.param_groups[0]['lr'], epoch) - self.writer.add_scalar('sick/dev/pearson_r', dev_scores[0], epoch) + self.writer.add_scalar('sick/dev/pearson_r', pearson, epoch) self.writer.add_scalar('sick/dev/kl_div_loss', new_loss, epoch) end = time.time() @@ -57,15 +57,15 @@ class SICKTrainer(Trainer): self.logger.info('Epoch {} finished in {:.2f} minutes'.format(epoch, duration / 60)) epoch_times.append(duration) - if dev_scores[0] > best_dev_score: - best_dev_score = dev_scores[0] - torch.save(self.model, self.model_outfile) + if pearson > best_dev_score: + best_dev_score = pearson + save_checkpoint(epoch, self.model.arch, self.model.state_dict(), self.optimizer.state_dict(), best_dev_score, self.model_outfile) if abs(prev_loss - new_loss) <= 0.0002: self.logger.info('Early stopping. Loss changed by less than 0.0002.') break prev_loss = new_loss - scheduler.step(dev_scores[0]) + scheduler.step(pearson) self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) diff --git a/common/trainers/trainer.py b/common/trainers/trainer.py index f16108a..23b74d2 100644 --- a/common/trainers/trainer.py +++ b/common/trainers/trainer.py @@ -4,8 +4,9 @@ class Trainer(object): Abstraction for training a model on a Dataset. """ - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): + def __init__(self, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): self.model = model + self.embedding = embedding self.optimizer = trainer_config['optimizer'] self.train_loader = train_loader self.batch_size = trainer_config['batch_size'] @@ -14,7 +15,6 @@ class Trainer(object): self.lr_reduce_factor = trainer_config['lr_reduce_factor'] self.patience = trainer_config['patience'] self.use_tensorboard = trainer_config['tensorboard'] - if self.use_tensorboard: from tensorboardX import SummaryWriter self.writer = SummaryWriter(log_dir=None, comment='' if trainer_config['run_label'] is None else trainer_config['run_label']) @@ -26,10 +26,16 @@ class Trainer(object): def evaluate(self, evaluator, dataset_name): scores, metric_names = evaluator.get_scores() + self.logger.info('Evaluation metrics for {}:'.format(dataset_name)) self.logger.info('\t'.join([' '] + metric_names)) self.logger.info('\t'.join([dataset_name] + list(map(str, scores)))) return scores + def get_sentence_embeddings(self, batch): + sent1 = self.embedding(batch.sentence_1).transpose(1, 2) + sent2 = self.embedding(batch.sentence_2).transpose(1, 2) + return sent1, sent2 + def train_epoch(self, epoch): raise NotImplementedError() diff --git a/common/trainers/trecqa_trainer.py b/common/trainers/trecqa_trainer.py index 03a4194..8828085 100644 --- a/common/trainers/trecqa_trainer.py +++ b/common/trainers/trecqa_trainer.py @@ -2,6 +2,4 @@ from .qa_trainer import QATrainer class TRECQATrainer(QATrainer): - - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): - super(TRECQATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + pass diff --git a/common/trainers/wikiqa_trainer.py b/common/trainers/wikiqa_trainer.py index dfc24a7..31ef3ee 100644 --- a/common/trainers/wikiqa_trainer.py +++ b/common/trainers/wikiqa_trainer.py @@ -2,6 +2,4 @@ from .qa_trainer import QATrainer class WikiQATrainer(QATrainer): - - def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): - super(WikiQATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + pass diff --git a/datasets/castor_dataset.py b/datasets/castor_dataset.py index 849d0fd..d192608 100644 --- a/datasets/castor_dataset.py +++ b/datasets/castor_dataset.py @@ -1,12 +1,11 @@ from abc import ABCMeta, abstractmethod import os -import numpy as np from sys import exit +import numpy as np +import torch from torchtext.data.dataset import Dataset from torchtext.data.example import Example -from torchtext.data.field import Field -import torch from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features @@ -20,6 +19,8 @@ class CastorPairDataset(Dataset, metaclass=ABCMeta): TEXT_FIELD = None EXT_FEATS_FIELD = None LABEL_FIELD = None + RAW_TEXT_FIELD = None + EXT_FEATS = 4 AID_FIELD = None @abstractmethod @@ -27,8 +28,9 @@ class CastorPairDataset(Dataset, metaclass=ABCMeta): """ Create a Castor dataset involving pairs of texts """ - 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)] + 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), ('sentence_1_raw', self.RAW_TEXT_FIELD), ('sentence_2_raw', self.RAW_TEXT_FIELD)] examples = [] with open(os.path.join(path, 'a.toks'), 'r') as f1, open(os.path.join(path, 'b.toks'), 'r') as f2: @@ -36,6 +38,7 @@ class CastorPairDataset(Dataset, metaclass=ABCMeta): 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) + self.word_to_doc_cnt = word_to_doc_cnt if not load_ext_feats: overlap_feats = get_pairwise_overlap_features(sent_list_1, sent_list_2, word_to_doc_cnt) @@ -46,7 +49,7 @@ class CastorPairDataset(Dataset, metaclass=ABCMeta): 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, i + 1] + example_list = [pair_id, l1, l2, ext_feats, label, i + 1, ' '.join(l1), ' '.join(l2)] example = Example.fromlist(example_list, fields) examples.append(example) diff --git a/datasets/msrvid.py b/datasets/msrvid.py index c7f389d..8a5e698 100644 --- a/datasets/msrvid.py +++ b/datasets/msrvid.py @@ -1,16 +1,13 @@ import math -import os import numpy as np import torch -from torchtext.data.example import Example -from torchtext.data.field import Field +from torchtext.data.field import Field, RawField from torchtext.data.iterator import BucketIterator from torchtext.data.pipeline import Pipeline from torchtext.vocab import Vectors from datasets.castor_dataset import CastorPairDataset -from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features def get_class_probs(sim, *args): @@ -35,6 +32,7 @@ class MSRVID(CastorPairDataset): 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) LABEL_FIELD = Field(sequential=False, tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True, postprocessing=Pipeline(get_class_probs)) + RAW_TEXT_FIELD = RawField() @staticmethod def sort_key(ex): diff --git a/datasets/sick.py b/datasets/sick.py index d9a0132..a78f686 100644 --- a/datasets/sick.py +++ b/datasets/sick.py @@ -1,16 +1,13 @@ import math -import os import numpy as np import torch -from torchtext.data.example import Example -from torchtext.data.field import Field +from torchtext.data.field import Field, RawField from torchtext.data.iterator import BucketIterator from torchtext.data.pipeline import Pipeline from torchtext.vocab import Vectors from datasets.castor_dataset import CastorPairDataset -from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features def get_class_probs(sim, *args): @@ -35,6 +32,7 @@ class SICK(CastorPairDataset): 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) LABEL_FIELD = Field(sequential=False, tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True, postprocessing=Pipeline(get_class_probs)) + RAW_TEXT_FIELD = RawField() @staticmethod def sort_key(ex): @@ -69,4 +67,5 @@ class SICK(CastorPairDataset): cls.TEXT_FIELD.build_vocab(train, val, test, vectors=vectors) - return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle, device=device) + return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle, + sort_within_batch=True, device=device) diff --git a/datasets/trecqa.py b/datasets/trecqa.py index 59183f8..f26ab67 100644 --- a/datasets/trecqa.py +++ b/datasets/trecqa.py @@ -1,14 +1,11 @@ import os import torch -from torchtext.data.field import Field +from torchtext.data.field import Field, RawField 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 TRECQA(CastorPairDataset): @@ -17,9 +14,9 @@ class TRECQA(CastorPairDataset): 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])) + EXT_FEATS_FIELD = Field(tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True, tokenize=lambda x: x) LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + RAW_TEXT_FIELD = RawField() VOCAB_SIZE = 0 @staticmethod @@ -37,7 +34,7 @@ class TRECQA(CastorPairDataset): return super(TRECQA, cls).splits(path, train=train, validation=validation, test=test, **kwargs) @classmethod - def iters(cls, path, vectors_name, vectors_dir, batch_size=64, shuffle=True, device=0, pt_file = False, vectors=None, unk_init=torch.Tensor.zero_): + def iters(cls, path, 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 vectors_name: name of word vectors file @@ -62,4 +59,4 @@ class TRECQA(CastorPairDataset): cls.VOCAB_SIZE = len(cls.TEXT_FIELD.vocab) - return BucketIterator.splits((train, validation, test), batch_size=batch_size, repeat=False, shuffle=shuffle, device=device) + return BucketIterator.splits((train, validation, test), batch_size=batch_size, repeat=False, shuffle=shuffle, sort_within_batch=True, device=device) diff --git a/datasets/wikiqa.py b/datasets/wikiqa.py index 5420f55..359a1dd 100644 --- a/datasets/wikiqa.py +++ b/datasets/wikiqa.py @@ -1,13 +1,11 @@ import os import torch -from torchtext.data.example import Example -from torchtext.data.field import Field +from torchtext.data.field import Field, RawField from torchtext.data.iterator import BucketIterator from torchtext.vocab import Vectors from datasets.castor_dataset import CastorPairDataset -from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features class WikiQA(CastorPairDataset): @@ -18,6 +16,8 @@ class WikiQA(CastorPairDataset): 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) LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + RAW_TEXT_FIELD = RawField() + VOCAB_SIZE = 0 @staticmethod def sort_key(ex): @@ -62,4 +62,4 @@ class WikiQA(CastorPairDataset): cls.VOCAB_SIZE = len(cls.TEXT_FIELD.vocab) return BucketIterator.splits((train, validation, test), batch_size=batch_size, repeat=False, shuffle=shuffle, - device=device) \ No newline at end of file + sort_within_batch=True, device=device) \ No newline at end of file diff --git a/mp_cnn/__main__.py b/mp_cnn/__main__.py index cf75446..060b11f 100644 --- a/mp_cnn/__main__.py +++ b/mp_cnn/__main__.py @@ -11,45 +11,11 @@ import torch.optim as optim from common.dataset import DatasetFactory from common.evaluation import EvaluatorFactory from common.train import TrainerFactory +from utils.serialization import load_checkpoint from .model import MPCNN -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('--word-vectors-dir', help='word vectors directory', default=os.path.join(os.pardir, os.pardir, 'Castor-data', 'embeddings', 'GloVe')) - parser.add_argument('--word-vectors-file', help='word vectors filename', default='glove.840B.300d.txt') - 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)') - parser.add_argument('--batch-size', type=int, default=64, help='input batch size for training (default: 64)') - parser.add_argument('--epochs', type=int, default=10, help='number of epochs to train (default: 10)') - parser.add_argument('--optimizer', type=str, default='adam', help='optimizer to use: adam or sgd (default: adam)') - parser.add_argument('--lr', type=float, default=0.001, help='learning rate (default: 0.001)') - parser.add_argument('--lr-reduce-factor', type=float, default=0.3, help='learning rate reduce factor after plateau (default: 0.3)') - parser.add_argument('--patience', type=float, default=2, help='learning rate patience after seeing plateau (default: 2)') - parser.add_argument('--momentum', type=float, default=0, help='momentum (default: 0)') - parser.add_argument('--epsilon', type=float, default=1e-8, help='Adam epsilon (default: 1e-8)') - parser.add_argument('--log-interval', type=int, default=10, help='how many batches to wait before logging training status (default: 10)') - parser.add_argument('--regularization', type=float, default=0.0001, help='Regularization for the optimizer (default: 0.0001)') - parser.add_argument('--max-window-size', type=int, default=3, help='windows sizes will be [1,max_window_size] and infinity (default: 300)') - parser.add_argument('--holistic-filters', type=int, default=300, help='number of holistic filters (default: 300)') - parser.add_argument('--per-dim-filters', type=int, default=20, help='number of per-dimension filters (default: 20)') - parser.add_argument('--hidden-units', type=int, default=150, help='number of hidden units in each of the two hidden layers (default: 150)') - parser.add_argument('--dropout', type=float, default=0.5, help='dropout probability (default: 0.5)') - parser.add_argument('--seed', type=int, default=1, help='random seed (default: 1)') - parser.add_argument('--tensorboard', action='store_true', default=False, help='use TensorBoard to visualize training (default: false)') - parser.add_argument('--run-label', type=str, help='label to describe run') - args = parser.parse_args() - - random.seed(args.seed) - np.random.seed(args.seed) - torch.manual_seed(args.seed) - if args.device != -1: - torch.cuda.manual_seed(args.seed) - - # logging setup +def get_logger(): logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) @@ -59,18 +25,82 @@ if __name__ == '__main__': ch.setFormatter(formatter) logger.addHandler(ch) + return logger + + +def evaluate_dataset(split_name, dataset_cls, model, embedding, loader, batch_size, device): + saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, loader, batch_size, device) + scores, metric_names = saved_model_evaluator.get_scores() + logger.info('Evaluation metrics for {}'.format(split_name)) + logger.info('\t'.join([' '] + metric_names)) + logger.info('\t'.join([split_name] + list(map(str, scores)))) + + +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('--word-vectors-dir', help='word vectors directory', + default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'GloVe')) + parser.add_argument('--word-vectors-file', help='word vectors filename', default='glove.840B.300d.txt') + parser.add_argument('--word-vectors-dim', type=int, default=300, + help='number of dimensions of word vectors (default: 300)') + 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('--wide-conv', action='store_true', default=False, + help='use wide convolution instead of narrow convolution (default: false)') + parser.add_argument('--attention', choices=['none', 'basic', 'idf'], default='none', help='type of attention to use') + parser.add_argument('--sparse-features', action='store_true', + default=False, help='use sparse features (default: false)') + parser.add_argument('--batch-size', type=int, default=64, help='input batch size for training (default: 64)') + parser.add_argument('--epochs', type=int, default=10, help='number of epochs to train (default: 10)') + parser.add_argument('--optimizer', type=str, default='adam', help='optimizer to use: adam or sgd (default: adam)') + parser.add_argument('--lr', type=float, default=0.001, help='learning rate (default: 0.001)') + parser.add_argument('--lr-reduce-factor', type=float, default=0.3, + help='learning rate reduce factor after plateau (default: 0.3)') + parser.add_argument('--patience', type=float, default=2, + help='learning rate patience after seeing plateau (default: 2)') + parser.add_argument('--momentum', type=float, default=0, help='momentum (default: 0)') + parser.add_argument('--epsilon', type=float, default=1e-8, help='Optimizer epsilon (default: 1e-8)') + parser.add_argument('--log-interval', type=int, default=10, + help='how many batches to wait before logging training status (default: 10)') + parser.add_argument('--regularization', type=float, default=0.0001, + help='Regularization for the optimizer (default: 0.0001)') + parser.add_argument('--max-window-size', type=int, default=3, + help='windows sizes will be [1,max_window_size] and infinity (default: 3)') + parser.add_argument('--holistic-filters', type=int, default=300, help='number of holistic filters (default: 300)') + parser.add_argument('--per-dim-filters', type=int, default=20, help='number of per-dimension filters (default: 20)') + parser.add_argument('--hidden-units', type=int, default=150, + help='number of hidden units in each of the two hidden layers (default: 150)') + parser.add_argument('--dropout', type=float, default=0.5, help='dropout probability (default: 0.5)') + parser.add_argument('--seed', type=int, default=1234, help='random seed (default: 1234)') + parser.add_argument('--tensorboard', action='store_true', default=False, + help='use TensorBoard to visualize training (default: false)') + parser.add_argument('--run-label', type=str, help='label to describe run') + args = parser.parse_args() + + device = torch.device(f'cuda:{args.device}' if torch.cuda.is_available() and args.device >= 0 else 'cpu') + + random.seed(args.seed) + np.random.seed(args.seed) + torch.manual_seed(args.seed) + if args.device != -1: + torch.cuda.manual_seed(args.seed) + + logger = get_logger() logger.info(pprint.pformat(vars(args))) dataset_cls, embedding, train_loader, test_loader, dev_loader \ = DatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device) filter_widths = list(range(1, args.max_window_size + 1)) + [np.inf] - model = MPCNN(embedding, args.holistic_filters, args.per_dim_filters, filter_widths, - args.hidden_units, dataset_cls.NUM_CLASSES, args.dropout, args.sparse_features) + ext_feats = dataset_cls.EXT_FEATS if args.sparse_features else 0 + model = MPCNN(args.word_vectors_dim, args.holistic_filters, args.per_dim_filters, filter_widths, + args.hidden_units, dataset_cls.NUM_CLASSES, args.dropout, ext_feats, + args.attention, args.wide_conv) - if args.device != -1: - with torch.cuda.device(args.device): - model.cuda() + model = model.to(device) + embedding = embedding.to(device) optimizer = None if args.optimizer == 'adam': @@ -80,9 +110,9 @@ if __name__ == '__main__': else: raise ValueError('optimizer not recognized: it should be either adam or sgd') - train_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, train_loader, args.batch_size, args.device) - test_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) - dev_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, dev_loader, args.batch_size, args.device) + train_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, train_loader, args.batch_size, args.device) + test_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, test_loader, args.batch_size, args.device) + dev_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, dev_loader, args.batch_size, args.device) trainer_config = { 'optimizer': optimizer, @@ -95,7 +125,7 @@ if __name__ == '__main__': 'run_label': args.run_label, 'logger': logger } - trainer = TrainerFactory.get_trainer(args.dataset, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + trainer = TrainerFactory.get_trainer(args.dataset, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) if not args.skip_training: total_params = 0 @@ -105,9 +135,12 @@ if __name__ == '__main__': logger.info('Total number of parameters: %s', total_params) trainer.train(args.epochs) - model = torch.load(args.model_outfile) - saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) - scores, metric_names = saved_model_evaluator.get_scores() - logger.info('Evaluation metrics for test') - logger.info('\t'.join([' '] + metric_names)) - logger.info('\t'.join(['test'] + list(map(str, scores)))) + _, _, state_dict, _, _ = load_checkpoint(args.model_outfile) + + for k, tensor in state_dict.items(): + state_dict[k] = tensor.to(device) + + model.load_state_dict(state_dict) + if dev_loader: + evaluate_dataset('dev', dataset_cls, model, embedding, dev_loader, args.batch_size, args.device) + evaluate_dataset('test', dataset_cls, model, embedding, test_loader, args.batch_size, args.device) diff --git a/mp_cnn/model.py b/mp_cnn/model.py index 520ad66..a6104d6 100644 --- a/mp_cnn/model.py +++ b/mp_cnn/model.py @@ -6,55 +6,89 @@ import torch.nn.functional as F class MPCNN(nn.Module): - def __init__(self, embedding, n_holistic_filters, n_per_dim_filters, filter_widths, hidden_layer_units, num_classes, dropout, ext_feats): + def __init__(self, n_word_dim, n_holistic_filters, n_per_dim_filters, filter_widths, hidden_layer_units, num_classes, dropout, ext_feats, attention, wide_conv): super(MPCNN, self).__init__() - self.embedding = embedding - self.n_word_dim = embedding.weight.size(1) + self.arch = 'mpcnn' + self.n_word_dim = n_word_dim self.n_holistic_filters = n_holistic_filters self.n_per_dim_filters = n_per_dim_filters self.filter_widths = filter_widths self.ext_feats = ext_feats - holistic_conv_layers = [] - per_dim_conv_layers = [] + self.attention = attention + self.wide_conv = wide_conv - for ws in filter_widths: + self.in_channels = n_word_dim if attention == 'none' else 2 * n_word_dim + + self._add_layers() + + # compute number of inputs to first hidden layer + n_feats = self._get_n_feats() + + self.final_layers = nn.Sequential( + nn.Linear(n_feats, hidden_layer_units), + nn.Tanh(), + nn.Dropout(dropout), + nn.Linear(hidden_layer_units, num_classes), + nn.LogSoftmax(1) + ) + + def _add_layers(self): + holistic_conv_layers_max = [] + holistic_conv_layers_min = [] + holistic_conv_layers_mean = [] + per_dim_conv_layers_max = [] + per_dim_conv_layers_min = [] + + for ws in self.filter_widths: if np.isinf(ws): continue - holistic_conv_layers.append(nn.Sequential( - nn.Conv1d(self.n_word_dim, n_holistic_filters, ws), + padding = ws-1 if self.wide_conv else 0 + + holistic_conv_layers_max.append(nn.Sequential( + nn.Conv1d(self.in_channels, self.n_holistic_filters, ws, padding=padding), nn.Tanh() )) - per_dim_conv_layers.append(nn.Sequential( - nn.Conv1d(self.n_word_dim, self.n_word_dim * n_per_dim_filters, ws, groups=self.n_word_dim), + holistic_conv_layers_min.append(nn.Sequential( + nn.Conv1d(self.in_channels, self.n_holistic_filters, ws, padding=padding), nn.Tanh() )) - self.holistic_conv_layers = nn.ModuleList(holistic_conv_layers) - self.per_dim_conv_layers = nn.ModuleList(per_dim_conv_layers) + holistic_conv_layers_mean.append(nn.Sequential( + nn.Conv1d(self.in_channels, self.n_holistic_filters, ws, padding=padding), + nn.Tanh() + )) - # compute number of inputs to first hidden layer - COMP_1_COMPONENTS_HOLISTIC, COMP_1_COMPONENTS_PER_DIM, COMP_2_COMPONENTS = 2 + n_holistic_filters, 2 + self.n_word_dim, 2 - EXT_FEATS = 4 if ext_feats else 0 - n_feat_h = 3 * len(self.filter_widths) * COMP_2_COMPONENTS - n_feat_v = ( + per_dim_conv_layers_max.append(nn.Sequential( + nn.Conv1d(self.in_channels, self.in_channels * self.n_per_dim_filters, ws, padding=padding, groups=self.in_channels), + nn.Tanh() + )) + + per_dim_conv_layers_min.append(nn.Sequential( + nn.Conv1d(self.in_channels, self.in_channels * self.n_per_dim_filters, ws, padding=padding, groups=self.in_channels), + nn.Tanh() + )) + + self.holistic_conv_layers_max = nn.ModuleList(holistic_conv_layers_max) + self.holistic_conv_layers_min = nn.ModuleList(holistic_conv_layers_min) + self.holistic_conv_layers_mean = nn.ModuleList(holistic_conv_layers_mean) + self.per_dim_conv_layers_max = nn.ModuleList(per_dim_conv_layers_max) + self.per_dim_conv_layers_min = nn.ModuleList(per_dim_conv_layers_min) + + def _get_n_feats(self): + COMP_1_COMPONENTS_HOLISTIC, COMP_1_COMPONENTS_PER_DIM, COMP_2_COMPONENTS = 2 + self.n_holistic_filters, 2 + self.in_channels, 2 + n_feats_h = 3 * self.n_holistic_filters * COMP_2_COMPONENTS + n_feats_v = ( # comparison units from holistic conv for min, max, mean pooling for non-infinite widths 3 * ((len(self.filter_widths) - 1) ** 2) * COMP_1_COMPONENTS_HOLISTIC + # comparison units from holistic conv for min, max, mean pooling for infinite widths 3 * 3 + # comparison units from per-dim conv - 2 * (len(self.filter_widths) - 1) * n_per_dim_filters * COMP_1_COMPONENTS_PER_DIM - ) - self.n_feat = n_feat_h + n_feat_v + EXT_FEATS - - self.final_layers = nn.Sequential( - nn.Linear(self.n_feat, hidden_layer_units), - nn.Tanh(), - nn.Dropout(dropout), - nn.Linear(hidden_layer_units, num_classes), - nn.LogSoftmax() + 2 * (len(self.filter_widths) - 1) * self.n_per_dim_filters * COMP_1_COMPONENTS_PER_DIM ) + n_feats = n_feats_h + n_feats_v + self.ext_feats + return n_feats def _get_blocks_for_sentence(self, sent): block_a = {} @@ -69,29 +103,48 @@ class MPCNN(nn.Module): } continue - holistic_conv_out = self.holistic_conv_layers[ws - 1](sent) + holistic_conv_out_max = self.holistic_conv_layers_max[ws - 1](sent) + holistic_conv_out_min = self.holistic_conv_layers_min[ws - 1](sent) + holistic_conv_out_mean = self.holistic_conv_layers_mean[ws - 1](sent) block_a[ws] = { - 'max': F.max_pool1d(holistic_conv_out, holistic_conv_out.size(2)).contiguous().view(-1, self.n_holistic_filters), - 'min': F.max_pool1d(-1 * holistic_conv_out, holistic_conv_out.size(2)).contiguous().view(-1, self.n_holistic_filters), - 'mean': F.avg_pool1d(holistic_conv_out, holistic_conv_out.size(2)).contiguous().view(-1, self.n_holistic_filters) + 'max': F.max_pool1d(holistic_conv_out_max, holistic_conv_out_max.size(2)).contiguous().view(-1, self.n_holistic_filters), + 'min': F.max_pool1d(-1 * holistic_conv_out_min, holistic_conv_out_min.size(2)).contiguous().view(-1, self.n_holistic_filters), + 'mean': F.avg_pool1d(holistic_conv_out_mean, holistic_conv_out_mean.size(2)).contiguous().view(-1, self.n_holistic_filters) } - per_dim_conv_out = self.per_dim_conv_layers[ws - 1](sent) + per_dim_conv_out_max = self.per_dim_conv_layers_max[ws - 1](sent) + per_dim_conv_out_min = self.per_dim_conv_layers_min[ws - 1](sent) block_b[ws] = { - 'max': F.max_pool1d(per_dim_conv_out, per_dim_conv_out.size(2)).contiguous().view(-1, self.n_word_dim, self.n_per_dim_filters), - 'min': F.max_pool1d(-1 * per_dim_conv_out, per_dim_conv_out.size(2)).contiguous().view(-1, self.n_word_dim, self.n_per_dim_filters) + 'max': F.max_pool1d(per_dim_conv_out_max, per_dim_conv_out_max.size(2)).contiguous().view(-1, self.in_channels, self.n_per_dim_filters), + 'min': F.max_pool1d(-1 * per_dim_conv_out_min, per_dim_conv_out_min.size(2)).contiguous().view(-1, self.in_channels, self.n_per_dim_filters) } return block_a, block_b def _algo_1_horiz_comp(self, sent1_block_a, sent2_block_a): comparison_feats = [] for pool in ('max', 'min', 'mean'): + regM1, regM2 = [], [] for ws in self.filter_widths: - x1 = sent1_block_a[ws][pool] - x2 = sent2_block_a[ws][pool] - batch_size = x1.size()[0] - comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) + x1 = sent1_block_a[ws][pool].unsqueeze(2) + x2 = sent2_block_a[ws][pool].unsqueeze(2) + if np.isinf(ws): + x1 = x1.expand(-1, self.n_holistic_filters, -1) + x2 = x2.expand(-1, self.n_holistic_filters, -1) + regM1.append(x1) + regM2.append(x2) + + regM1 = torch.cat(regM1, dim=2) + regM2 = torch.cat(regM2, dim=2) + + # Cosine similarity + comparison_feats.append(F.cosine_similarity(regM1, regM2, dim=2)) + # Euclidean distance + pairwise_distances = [] + for x1, x2 in zip(regM1, regM2): + dist = F.pairwise_distance(x1, x2).view(1, -1) + pairwise_distances.append(dist) + comparison_feats.append(torch.cat(pairwise_distances)) + return torch.cat(comparison_feats, dim=1) def _algo_2_vert_comp(self, sent1_block_a, sent2_block_a, sent1_block_b, sent2_block_b): @@ -100,12 +153,11 @@ class MPCNN(nn.Module): for pool in ('max', 'min', 'mean'): for ws1 in self.filter_widths: x1 = sent1_block_a[ws1][pool] - batch_size = x1.size()[0] for ws2 in self.filter_widths: x2 = sent2_block_a[ws2][pool] if (not np.isinf(ws1) and not np.isinf(ws2)) or (np.isinf(ws1) and np.isinf(ws2)): - comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) + comparison_feats.append(F.cosine_similarity(x1, x2).unsqueeze(1)) + comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(1)) comparison_feats.append(torch.abs(x1 - x2)) for pool in ('max', 'min'): @@ -115,19 +167,61 @@ class MPCNN(nn.Module): for i in range(0, self.n_per_dim_filters): x1 = oG_1B[:, :, i] x2 = oG_2B[:, :, i] - batch_size = x1.size()[0] - comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) + comparison_feats.append(F.cosine_similarity(x1, x2).unsqueeze(1)) + comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(1)) comparison_feats.append(torch.abs(x1 - x2)) return torch.cat(comparison_feats, dim=1) - def forward(self, sent1_idx, sent2_idx, ext_feats=None): - # Select embedding - sent1 = self.embedding(sent1_idx).transpose(1, 2) - sent2 = self.embedding(sent2_idx).transpose(1, 2) + def concat_attention(self, sent1, sent2, word_to_doc_count=None, raw_sent1=None, raw_sent2=None): + sent1_transposed = sent1.transpose(1, 2) + attention_dot = torch.bmm(sent1_transposed, sent2) + sent1_norms = torch.norm(sent1_transposed, p=2, dim=2, keepdim=True) + sent2_norms = torch.norm(sent2, p=2, dim=1, keepdim=True) + attention_norms = torch.bmm(sent1_norms, sent2_norms) + attention_matrix = attention_dot / attention_norms - # Sentence modeling module + if self.attention == 'idf' and word_to_doc_count is not None: + idf_matrix1 = sent1.data.new_ones(sent1.size(0), sent1.size(2)) + for i, sent in enumerate(raw_sent1): + for j, word in enumerate(sent.split(' ')): + idf_matrix1[i, j] /= word_to_doc_count.get(word, 1) + + idf_matrix2 = sent2.data.new_ones(sent2.size(0), sent2.size(2)).fill_(1) + for i, sent in enumerate(raw_sent2): + for j, word in enumerate(sent.split(' ')): + idf_matrix2[i, j] /= word_to_doc_count.get(word, 1) + + sum_row = (attention_matrix * idf_matrix2.unsqueeze(1)).sum(2) + sum_col = (attention_matrix * idf_matrix1.unsqueeze(2)).sum(1) + else: + sum_row = attention_matrix.sum(2) + sum_col = attention_matrix.sum(1) + + if self.attention == 'idf' and word_to_doc_count is not None: + for i, sent in enumerate(raw_sent1): + for j, word in enumerate(sent.split(' ')): + sum_row[i, j] /= word_to_doc_count.get(word, 1) + + for i, sent in enumerate(raw_sent2): + for j, word in enumerate(sent.split(' ')): + sum_col[i, j] /= word_to_doc_count.get(word, 1) + + attention_weight_vec1 = F.softmax(sum_row, 1) + attention_weight_vec2 = F.softmax(sum_col, 1) + + attention_weighted_sent1 = attention_weight_vec1.unsqueeze(1).expand(-1, self.n_word_dim, -1) * sent1 + attention_weighted_sent2 = attention_weight_vec2.unsqueeze(1).expand(-1, self.n_word_dim, -1) * sent2 + attention_emb1 = torch.cat((attention_weighted_sent1, sent1), dim=1) + attention_emb2 = torch.cat((attention_weighted_sent2, sent2), dim=1) + return attention_emb1, attention_emb2 + + def forward(self, sent1, sent2, ext_feats=None, word_to_doc_count=None, raw_sent1=None, raw_sent2=None, sent1_nonstatic=None, sent2_nonstatic=None): + # Attention + if self.attention != 'none': + sent1, sent2 = self.concat_attention(sent1, sent2, word_to_doc_count, raw_sent1, raw_sent2) + + # Sentence modelling module sent1_block_a, sent1_block_b = self._get_blocks_for_sentence(sent1) sent2_block_a, sent2_block_b = self._get_blocks_for_sentence(sent2) diff --git a/nce/nce_pairwise_mp/evaluators/qa_evaluator.py b/nce/nce_pairwise_mp/evaluators/qa_evaluator.py index c546181..918f191 100644 --- a/nce/nce_pairwise_mp/evaluators/qa_evaluator.py +++ b/nce/nce_pairwise_mp/evaluators/qa_evaluator.py @@ -7,7 +7,7 @@ from utils.relevancy_metrics import get_map_mrr class QAEvaluator(Evaluator): def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(QAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) + super(QAEvaluator, self).__init__(dataset_cls, model, None, data_loader, batch_size, device) def get_scores(self): self.model.eval() diff --git a/nce/nce_pairwise_mp/trainers/qa_trainer.py b/nce/nce_pairwise_mp/trainers/qa_trainer.py index b935898..e3b8ec8 100644 --- a/nce/nce_pairwise_mp/trainers/qa_trainer.py +++ b/nce/nce_pairwise_mp/trainers/qa_trainer.py @@ -11,7 +11,7 @@ from utils.nce_neighbors import get_nearest_neg_id, get_random_neg_id, get_batch class QATrainer(Trainer): def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None, weighting=False): - super(QATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + super(QATrainer, self).__init__(model, None, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) self.loss = torch.nn.MarginRankingLoss(margin=1, size_average=True) self.question2answer = {} self.best_dev_map = 0 diff --git a/utils/serialization.py b/utils/serialization.py new file mode 100644 index 0000000..51b06b6 --- /dev/null +++ b/utils/serialization.py @@ -0,0 +1,23 @@ +""" +Utils for serialization +""" +import torch + + +def save_checkpoint(epoch, arch, state_dict, optimizer_state, eval_metric, filename): + for k, tensor in state_dict.items(): + state_dict[k] = tensor.cpu() + + state = { + 'epoch': epoch, + 'arch': arch, + 'state_dict': state_dict, + 'optimizer_state': None, # currently do not save optimizer state + 'eval_metric': eval_metric + } + torch.save(state, filename) + + +def load_checkpoint(filename): + state = torch.load(filename) + return state['epoch'], state['arch'], state['state_dict'], state['optimizer_state'], state['eval_metric']