Refactor VDPWI to use common API (#109)

This commit is contained in:
Ralph Tang
2018-05-25 01:30:53 -04:00
committed by GitHub
parent 0c3a91c443
commit 494ce36575
4 changed files with 178 additions and 168 deletions
+8 -2
View File
@@ -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))
+2
View File
@@ -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'])
+121 -119
View File
@@ -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()
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)
+47 -47
View File
@@ -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