diff --git a/common/trainers/sick_trainer.py b/common/trainers/sick_trainer.py index ca50e24..0b3858b 100644 --- a/common/trainers/sick_trainer.py +++ b/common/trainers/sick_trainer.py @@ -1,5 +1,6 @@ import time +import torch.nn as nn import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau @@ -22,6 +23,8 @@ class SICKTrainer(Trainer): loss = F.kl_div(output, batch.label, size_average=False) total_loss += loss.item() loss.backward() + if self.clip_norm: + nn.utils.clip_grad_norm(self.model.parameters(), self.clip_norm) self.optimizer.step() if batch_idx % self.log_interval == 0: self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( @@ -36,7 +39,9 @@ class SICKTrainer(Trainer): return total_loss def train(self, epochs): - scheduler = ReduceLROnPlateau(self.optimizer, mode='max', factor=self.lr_reduce_factor, patience=self.patience) + scheduler = None + if self.lr_reduce_factor != 1 and self.lr_reduce_factor != None: + scheduler = ReduceLROnPlateau(self.optimizer, mode='max', factor=self.lr_reduce_factor, patience=self.patience) epoch_times = [] prev_loss = -1 best_dev_score = -1 @@ -66,6 +71,7 @@ class SICKTrainer(Trainer): break prev_loss = new_loss - scheduler.step(pearson) + if scheduler is not None: + 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 23b74d2..553824e 100644 --- a/common/trainers/trainer.py +++ b/common/trainers/trainer.py @@ -15,6 +15,8 @@ class Trainer(object): self.lr_reduce_factor = trainer_config['lr_reduce_factor'] self.patience = trainer_config['patience'] self.use_tensorboard = trainer_config['tensorboard'] + self.clip_norm = trainer_config.get('clip_norm') + 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']) diff --git a/vdpwi/__main__.py b/vdpwi/__main__.py index 05a56be..7b3494e 100644 --- a/vdpwi/__main__.py +++ b/vdpwi/__main__.py @@ -1,139 +1,141 @@ -from collections import namedtuple +import argparse +import logging +import os +import pprint +import random -from tqdm import tqdm import numpy as np -import scipy.stats as stats import torch import torch.optim as optim -import torch.nn as nn -import torch.nn.functional as F -import torch.utils as utils -from utils.log import LogWriter -import data -import model as mod +from common.dataset import DatasetFactory +from common.evaluation import EvaluatorFactory +from common.train import TrainerFactory +from utils.serialization import load_checkpoint +from .model import VDPWIModel -Context = namedtuple("Context", "model, train_loader, dev_loader, test_loader, optimizer, criterion, params, log_writer") -EvaluateResult = namedtuple("EvaluateResult", "pearsonr, spearmanr") -def create_context(config): - def collate_fn(batch): - emb1 = [] - emb2 = [] - labels = [] - cmp_labels = [] - pad_cube = [] - max_len1 = 0; max_len2 = 0 +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)))) - for s1, s2, l, cl in batch: - emb1.append(s1) - emb2.append(s2) - max_len1 = max(max_len1, len(s1)) - max_len2 = max(max_len2, len(s2)) - labels.append(l) - cmp_labels.append(cl) +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='PyTorch implementation of VDPWI') + 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('--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('--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=5E-4, help='learning rate (default: 0.001)') + parser.add_argument('--lr-reduce-factor', type=float, default=1, 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.1, help='momentum (default: 0.1)') + 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=1E-5, help='Regularization for the optimizer (default: 0.00001)') + parser.add_argument('--hidden-units', type=int, default=150, help='number of hidden units in the RNN') + 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') + # VDPWI args + parser.add_argument('--classifier', type=str, default='vdpwi', choices=['vdpwi', 'resnet']) + parser.add_argument('--clip-norm', type=float, default=50) + parser.add_argument('--decay', type=float, default=0.95) + parser.add_argument('--res-fmaps', type=int, default=32) + parser.add_argument('--res-layers', type=int, default=16) + parser.add_argument('--rnn-hidden-dim', type=int, default=250) + args = parser.parse_args() - for s1, s2 in zip(emb1, emb2): - pad1 = (max_len1 - len(s1)) - pad2 = (max_len2 - len(s2)) - pad_mask = np.ones((max_len1, max_len2)) - pad_mask[:len(s1), :len(s2)] = 0 - pad_cube.append(pad_mask) - s1.extend([embedding.weight.size(0) - 1] * pad1) - s2.extend([embedding.weight.size(0) - 1] * pad2) + device = torch.device(f'cuda:{args.device}' if torch.cuda.is_available() and args.device >= 0 else 'cpu') - pad_cube = np.array(pad_cube) - emb1 = torch.LongTensor(emb1) - emb2 = torch.LongTensor(emb2) - labels = torch.Tensor(labels) - emb1 = torch.autograd.Variable(emb1, requires_grad=False) - emb2 = torch.autograd.Variable(emb2, requires_grad=False) - labels = torch.autograd.Variable(labels, requires_grad=False) - pad_cube = torch.autograd.Variable(torch.from_numpy(pad_cube).float(), requires_grad=False) - if not config.cpu: - emb1 = emb1.cuda() - emb2 = emb2.cuda() - labels = labels.cuda() - pad_cube = pad_cube.cuda() - return emb1, emb2, labels, pad_cube, cmp_labels + random.seed(args.seed) + np.random.seed(args.seed) + torch.manual_seed(args.seed) + if args.device != -1: + torch.cuda.manual_seed(args.seed) - embedding, (train_set, dev_set, test_set) = data.load_dataset(config.dataset) - model = mod.VDPWIModel(embedding, config) - if config.restore: - model.load(config.input_file) - if not config.cpu: - model = model.cuda() + # logging setup + logger = logging.getLogger(__name__) + logger.setLevel(logging.INFO) - train_loader = utils.data.DataLoader(train_set, shuffle=True, batch_size=config.mbatch_size, collate_fn=collate_fn) - dev_loader = utils.data.DataLoader(dev_set, batch_size=1, collate_fn=collate_fn) - test_loader = utils.data.DataLoader(test_set, batch_size=1, collate_fn=collate_fn) + ch = logging.StreamHandler() + ch.setLevel(logging.DEBUG) + formatter = logging.Formatter('%(levelname)s - %(message)s') + ch.setFormatter(formatter) + logger.addHandler(ch) - params = list(filter(lambda x: x.requires_grad, model.parameters())) - if config.optimizer == "adam": - optimizer = optim.Adam(params, lr=config.lr, weight_decay=config.weight_decay) - elif config.optimizer == "sgd": - optimizer = optim.SGD(params, lr=config.lr, momentum=config.momentum, weight_decay=config.weight_decay) - elif config.optimizer == "rmsprop": - optimizer = optim.RMSprop(params, lr=config.lr, alpha=config.decay, momentum=config.momentum, weight_decay=config.weight_decay) - criterion = nn.KLDivLoss() - log_writer = LogWriter() - return Context(model, train_loader, dev_loader, test_loader, optimizer, criterion, params, log_writer) + logger.info(pprint.pformat(vars(args))) -def test(config): - context = create_context(config) - result = evaluate(context, context.test_loader) - print("Final test result: {}".format(result)) + 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) -def evaluate(context, data_loader): - model = context.model - model.eval() - predictions = [] - true_labels = [] - for sent1, sent2, _, pad_cube, truth in data_loader: - scores = model(sent1, sent2, pad_cube) - scores = F.softmax(scores).cpu().data.numpy()[0] - prediction = np.dot(np.arange(1, len(scores) + 1), scores) - predictions.append(prediction); true_labels.append(truth[0][0]) - - pearsonr = stats.pearsonr(predictions, true_labels)[0] - spearmanr = stats.spearmanr(predictions, true_labels)[0] - context.log_writer.log_dev_metrics(pearsonr, spearmanr) - return EvaluateResult(pearsonr, spearmanr) + model_config = { + 'classifier': args.classifier, + 'rnn_hidden_dim': args.rnn_hidden_dim, + 'n_labels': dataset_cls.NUM_CLASSES, + 'device': device, + 'res_layers': args.res_layers, + 'res_fmaps': args.res_fmaps + } -def train(config): - context = create_context(config) - context.log_writer.log_hyperparams() - best_dev_pr = 0 - for epoch_no in range(config.n_epochs): - print("Epoch number: {}".format(epoch_no + 1)) - loader_wrapper = tqdm(context.train_loader, total=len(context.train_loader), desc="Loss") - context.model.train() - loss = 0 - for sent1, sent2, label_pmf, pad_cube, _ in loader_wrapper: - context.optimizer.zero_grad() - scores = F.log_softmax(context.model(sent1, sent2, pad_cube)) + model = VDPWIModel(args.word_vectors_dim, model_config) + model.to(device) + embedding = embedding.to(device) - loss = context.criterion(scores, label_pmf) - loss.backward() - nn.utils.clip_grad_norm(context.params, config.clip_norm) - context.optimizer.step() + optimizer = None + if args.optimizer == 'adam': + optimizer = optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.regularization, eps=args.epsilon) + elif args.optimizer == 'sgd': + optimizer = optim.SGD(model.parameters(), lr=args.lr, momentum=args.momentum, weight_decay=args.regularization) + elif args.optimizer == "rmsprop": + optimizer = optim.RMSprop(model.parameters(), lr=args.lr, momentum=args.momentum, alpha=config.decay, + weight_decay=args.regularization) + else: + raise ValueError('optimizer not recognized: it should be one of adam, sgd, or rmsprop') - loss = loss.cpu().data[0] - loader_wrapper.set_description("Loss: {:<8}".format(round(loss, 5))) - context.log_writer.log_train_loss(loss) - result = evaluate(context, context.dev_loader) - print("Dev result: {}".format(result)) - if best_dev_pr < result.pearsonr: - best_dev_pr = result.pearsonr - print("Saving best model...") - context.model.save(config.output_file) + 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) -def main(): - config = data.Configs.base_config() - if config.mode == "train": - train(config) - elif config.mode == "test": - test(config) + trainer_config = { + 'optimizer': optimizer, + 'batch_size': args.batch_size, + 'log_interval': args.log_interval, + 'model_outfile': args.model_outfile, + 'lr_reduce_factor': args.lr_reduce_factor, + 'patience': args.patience, + 'tensorboard': args.tensorboard, + 'run_label': args.run_label, + 'logger': logger, + 'clip_norm': args.clip_norm + } -if __name__ == "__main__": - main() \ No newline at end of file + 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 + for param in model.parameters(): + size = [s for s in param.size()] + total_params += np.prod(size) + logger.info('Total number of parameters: %s', total_params) + trainer.train(args.epochs) + + _, _, 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/vdpwi/model.py b/vdpwi/model.py index c31bcbc..32b4d1f 100644 --- a/vdpwi/model.py +++ b/vdpwi/model.py @@ -1,20 +1,8 @@ -from torch.autograd import Variable import torch import torch.nn as nn import torch.nn.functional as F -import torchvision.models as models import numpy as np -class SerializableModule(nn.Module): - def __init__(self): - super().__init__() - - def save(self, filename): - torch.save(self.state_dict(), filename) - - def load(self, filename): - self.load_state_dict(torch.load(filename, map_location=lambda storage, loc: storage)) - def hard_pad2d(x, pad): def pad_side(idx): pad_len = max(pad - x.size(idx), 0) @@ -24,18 +12,16 @@ def hard_pad2d(x, pad): x = F.pad(x, padding) return x[:, :, :pad, :pad] -class ResNet(SerializableModule): +class ResNet(nn.Module): def __init__(self, config): super().__init__() - n_layers = config.res_layers - n_maps = config.res_fmaps - n_labels = config.n_labels + n_layers = config['res_layers'] + n_maps = config['res_fmaps'] + n_labels = config['n_labels'] self.conv0 = nn.Conv2d(12, n_maps, (3, 3), padding=1) - self.convs = [nn.Conv2d(n_maps, n_maps, (3, 3), padding=1) for _ in range(n_layers)] + self.convs = nn.ModuleList([nn.Conv2d(n_maps, n_maps, (3, 3), padding=1) for _ in range(n_layers)]) self.output = nn.Linear(n_maps, n_labels) self.input_len = None - for i, conv in enumerate(self.convs): - self.add_module("conv{}".format(i + 1), conv) def forward(self, x): x = F.relu(self.conv0(x)) @@ -48,7 +34,7 @@ class ResNet(SerializableModule): x = torch.mean(x.view(x.size(0), x.size(1), -1), 2) return self.output(x) -class VDPWIConvNet(SerializableModule): +class VDPWIConvNet(nn.Module): def __init__(self, config): super().__init__() def make_conv(n_in, n_out): @@ -63,7 +49,7 @@ class VDPWIConvNet(SerializableModule): self.conv5 = make_conv(192, 128) self.maxpool2 = nn.MaxPool2d(2, ceil_mode=True) self.dnn = nn.Linear(128, 128) - self.output = nn.Linear(128, config.n_labels) + self.output = nn.Linear(128, config['n_labels']) self.input_len = 32 def forward(self, x): @@ -75,20 +61,35 @@ class VDPWIConvNet(SerializableModule): x = self.maxpool2(F.relu(self.conv4(x))) x = pool_final(F.relu(self.conv5(x))) x = F.relu(self.dnn(x.view(x.size(0), -1))) - return self.output(x) + return F.log_softmax(self.output(x), 1) -class VDPWIModel(SerializableModule): - def __init__(self, embedding, config): +class VDPWIModel(nn.Module): + def __init__(self, dim, config): super().__init__() - self.hidden_dim = config.rnn_hidden_dim - self.rnn = nn.LSTM(300, self.hidden_dim, 1, batch_first=True) - self.embedding = embedding - self.use_cuda = not config.cpu - if config.classifier == "vdpwi": + self.arch = 'vdpwi' + self.hidden_dim = config['rnn_hidden_dim'] + self.rnn = nn.LSTM(dim, self.hidden_dim, 1, batch_first=True) + self.device = config['device'] + if config['classifier'] == 'vdpwi': self.classifier_net = VDPWIConvNet(config) - elif config.classifier == "resnet": + elif config['classifier'] == 'resnet': self.classifier_net = ResNet(config) + def create_pad_cube(self, sent1, sent2): + pad_cube = [] + max_len1 = max([len(s.split()) for s in sent1]) + max_len2 = max([len(s.split()) for s in sent2]) + + for s1, s2 in zip(sent1, sent2): + pad1 = (max_len1 - len(s1.split())) + pad2 = (max_len2 - len(s2.split())) + pad_mask = np.ones((max_len1, max_len2)) + pad_mask[:len(s1), :len(s2)] = 0 + pad_cube.append(pad_mask) + + pad_cube = np.array(pad_cube) + return torch.from_numpy(pad_cube).float().to(self.device).unsqueeze(0) + def compute_sim_cube(self, seq1, seq2): def compute_sim(prism1, prism2): prism1_len = prism1.norm(dim=3) @@ -97,7 +98,7 @@ class VDPWIModel(SerializableModule): dot_prod = torch.matmul(prism1.unsqueeze(3), prism2.unsqueeze(4)) dot_prod = dot_prod.squeeze(3).squeeze(3) cos_dist = dot_prod / (prism1_len * prism2_len + 1E-8) - l2_dist = -((prism1 - prism2).norm(dim=3)) + l2_dist = ((prism1 - prism2).norm(dim=3)) return torch.stack([dot_prod, cos_dist, l2_dist], 1) def compute_prism(seq1, seq2): @@ -107,9 +108,8 @@ class VDPWIModel(SerializableModule): prism2 = prism2.permute(1, 0, 2, 3).contiguous() return compute_sim(prism1, prism2) - sim_cube = Variable(torch.Tensor(seq1.size(0), 12, seq1.size(1), seq2.size(1))) - if self.use_cuda: - sim_cube = sim_cube.cuda() + sim_cube = torch.Tensor(seq1.size(0), 12, seq1.size(1), seq2.size(1)) + sim_cube = sim_cube.to(self.device) seq1_f = seq1[:, :, :self.hidden_dim] seq1_b = seq1[:, :, self.hidden_dim:] seq2_f = seq2[:, :, :self.hidden_dim] @@ -125,9 +125,7 @@ class VDPWIModel(SerializableModule): pad_cube = pad_cube.repeat(12, 1, 1, 1) pad_cube = pad_cube.permute(1, 0, 2, 3).contiguous() sim_cube = neg_magic * pad_cube + sim_cube - mask = Variable(torch.Tensor(*sim_cube.size())) - if self.use_cuda: - mask = mask.cuda() + mask = torch.Tensor(*sim_cube.size()).to(self.device) mask[:, :, :, :] = 0.1 def build_mask(index): @@ -149,20 +147,22 @@ class VDPWIModel(SerializableModule): focus_cube = mask * sim_cube * (1 - pad_cube) return focus_cube - def forward(self, x1, x2, pad_cube): - x1 = self.embedding(x1) - x2 = self.embedding(x2) - seq1f, _ = self.rnn(x1) - seq2f, _ = self.rnn(x2) - seq1b, _ = self.rnn(torch.cat(x1.split(1, 1)[::-1], 1)) - seq2b, _ = self.rnn(torch.cat(x2.split(1, 1)[::-1], 1)) + def forward(self, sent1, sent2, ext_feats=None, word_to_doc_count=None, raw_sent1=None, raw_sent2=None): + pad_cube = self.create_pad_cube(raw_sent1, raw_sent2) + sent1 = sent1.permute(0, 2, 1).contiguous() + sent2 = sent2.permute(0, 2, 1).contiguous() + seq1f, _ = self.rnn(sent1) + seq2f, _ = self.rnn(sent2) + seq1b, _ = self.rnn(torch.cat(sent1.split(1, 1)[::-1], 1)) + seq2b, _ = self.rnn(torch.cat(sent2.split(1, 1)[::-1], 1)) seq1 = torch.cat([seq1f, seq1b], 2) seq2 = torch.cat([seq2f, seq2b], 2) sim_cube = self.compute_sim_cube(seq1, seq2) truncate = self.classifier_net.input_len + sim_cube = sim_cube[:, :, :pad_cube.size(2), :pad_cube.size(3)].contiguous() if truncate is not None: sim_cube = sim_cube[:, :, :truncate, :truncate].contiguous() - pad_cube = pad_cube[:, :truncate, :truncate].contiguous() + pad_cube = pad_cube[:, :, :sim_cube.size(2), :sim_cube.size(3)].contiguous() focus_cube = self.compute_focus_cube(sim_cube, pad_cube) - logits = self.classifier_net(focus_cube) - return logits + log_prob = self.classifier_net(focus_cube) + return log_prob