mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
MP-CNN with Bugs Fixed and PyTorch v0.4 (#107)
* Refactor datasets * Update evaluators * Update trainers * Update main and MP-CNN model * Add serialization util * Fix bugs * Refactoring for NCE to use new parent class
This commit is contained in:
+5
-13
@@ -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))
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+2
-2
@@ -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
|
||||
)
|
||||
|
||||
@@ -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))
|
||||
@@ -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))
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
+2
-4
@@ -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):
|
||||
|
||||
+4
-5
@@ -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)
|
||||
|
||||
+5
-8
@@ -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)
|
||||
|
||||
+4
-4
@@ -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)
|
||||
sort_within_batch=True, device=device)
|
||||
+84
-51
@@ -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)
|
||||
|
||||
+144
-50
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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']
|
||||
Reference in New Issue
Block a user