mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Refactor VDPWI to use common API (#109)
This commit is contained in:
@@ -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))
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user