diff --git a/.gitignore b/.gitignore index 31104a5..96fe215 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ *.sublime-workspace local_* +runs # Byte-compiled / optimized / DLL files __pycache__/ diff --git a/vdpwi/__main__.py b/vdpwi/__main__.py index f343d44..99a8e03 100644 --- a/vdpwi/__main__.py +++ b/vdpwi/__main__.py @@ -9,10 +9,11 @@ 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 -Context = namedtuple("Context", "model, train_loader, dev_loader, test_loader, optimizer, criterion, params") +Context = namedtuple("Context", "model, train_loader, dev_loader, test_loader, optimizer, criterion, params, log_writer") EvaluateResult = namedtuple("EvaluateResult", "pearsonr, spearmanr") def create_context(config): @@ -20,10 +21,12 @@ def create_context(config): emb1 = [] emb2 = [] labels = [] - for s1, s2, l in batch: + cmp_labels = [] + for s1, s2, l, cl in batch: emb1.append(s1) emb2.append(s2) labels.append(l) + cmp_labels.append(cl) emb1 = torch.LongTensor(emb1) emb2 = torch.LongTensor(emb2) labels = torch.Tensor(labels) @@ -34,7 +37,7 @@ def create_context(config): emb1 = emb1.cuda() emb2 = emb2.cuda() labels = labels.cuda() - return emb1, emb2, labels + return emb1, emb2, labels, cmp_labels embedding, (train_set, dev_set, test_set) = data.load_dataset(config.dataset) model = mod.VDPWIModel(embedding, config) @@ -48,42 +51,65 @@ def create_context(config): test_loader = utils.data.DataLoader(test_set, batch_size=1, collate_fn=collate_fn) params = list(filter(lambda x: x.requires_grad, model.parameters())) - optimizer = optim.RMSprop(params, lr=config.lr, alpha=config.decay, momentum=config.momentum) + optimizer = optim.Adam(params, lr=config.lr, weight_decay=config.weight_decay) + # optimizer = optim.SGD(params, lr=config.lr, momentum=config.momentum, weight_decay=config.weight_decay) criterion = nn.KLDivLoss() - return Context(model, train_loader, dev_loader, test_loader, optimizer, criterion, params) + log_writer = LogWriter() + return Context(model, train_loader, dev_loader, test_loader, optimizer, criterion, params, log_writer) def test(config): - pass + context = create_context(config) + result = evaluate(context, context.test_loader) + print("Final test result: {}".format(result)) -def evaluate(model, data_loader): +def evaluate(context, data_loader): + model = context.model model.eval() predictions = [] true_labels = [] - for sent1, sent2, label_pmf in data_loader: + for sent1, sent2, _, truth in data_loader: scores = model(sent1, sent2) scores = F.softmax(scores).cpu().data.numpy()[0] prediction = np.dot(np.arange(1, len(scores) + 1), scores) - truth = np.dot(np.arange(1, len(scores) + 1), label_pmf.cpu().data.numpy()[0]) - predictions.append(prediction); true_labels.append(truth) - return EvaluateResult(stats.pearsonr(predictions, true_labels)[0], stats.spearmanr(predictions, true_labels)[0]) + 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) 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(enumerate(context.train_loader), total=len(context.train_loader), desc="Loss") context.model.train() - for i, (sent1, sent2, label_pmf) in loader_wrapper: + loss = 0 + for i, (sent1, sent2, label_pmf, _) in loader_wrapper: context.optimizer.zero_grad() scores = F.log_softmax(context.model(sent1, sent2)) - loss = context.criterion(scores, label_pmf) - loss.backward() - nn.utils.clip_grad_norm(context.params, 50) - loader_wrapper.set_description("Loss = {}".format(loss.cpu().data[0])) - context.optimizer.step() - result = evaluate(context.model, context.dev_loader) - print(result) + loss = context.criterion(scores, label_pmf) + loss + if i % config.mbatch_size == (config.mbatch_size - 1): + loss /= config.mbatch_size + loss.backward() + nn.utils.clip_grad_norm(context.params, 5) + context.optimizer.step() + + loss = loss.cpu().data[0] + loader_wrapper.set_description("Loss: {:<8}".format(round(loss, 5))) + context.log_writer.log_train_loss(loss) + loss = 0 + 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) + test_result = evaluate(context, context.test_loader) + print("Final test result: {}".format(test_result)) def main(): config = data.Configs.base_config() diff --git a/vdpwi/data.py b/vdpwi/data.py index a67aa8f..4180f03 100644 --- a/vdpwi/data.py +++ b/vdpwi/data.py @@ -12,16 +12,18 @@ class Configs(object): parser.add_argument("--cpu", action="store_true", default=False) parser.add_argument("--dataset", type=str, default="sick", choices=["sick"]) parser.add_argument("--decay", type=float, default=0.95) - parser.add_argument("--input_model", type=str, default="local_saves/model.pt") - parser.add_argument("--lr", type=float, default=1E-4) - parser.add_argument("--mbatch_size", type=int, default=1) + parser.add_argument("--input_file", type=str, default="local_saves/model.pt") + parser.add_argument("--lr", type=float, default=1E-3) + parser.add_argument("--mbatch_size", type=int, default=16) parser.add_argument("--mode", type=str, default="train", choices=["train", "test"]) parser.add_argument("--momentum", type=float, default=0.9) parser.add_argument("--n_epochs", type=int, default=40) parser.add_argument("--n_labels", type=int, default=5) - parser.add_argument("--output_model", type=str, default="local_saves/model.pt") + parser.add_argument("--optimizer", type=str, default="adam", choices=["adam", "sgd", "rmsprop"]) + parser.add_argument("--output_file", type=str, default="local_saves/model.pt") parser.add_argument("--restore", action="store_true", default=False) parser.add_argument("--rnn_hidden_dim", type=int, default=250) + parser.add_argument("--weight_decay", type=float, default=5E-4) parser.add_argument("--wordvecs_file", type=str, default="local_data/glove/glove.840B.300d.txt") return parser.parse_known_args()[0] @@ -34,14 +36,16 @@ class Configs(object): return parser.parse_known_args()[0] class LabeledEmbeddedDataset(data.Dataset): - def __init__(self, sentence_indices1, sentence_indices2, labels): + def __init__(self, sentence_indices1, sentence_indices2, labels, compare_labels=None): assert len(sentence_indices1) == len(labels) == len(sentence_indices2) self.sentence_indices1 = sentence_indices1 self.sentence_indices2 = sentence_indices2 self.labels = labels + self.compare_labels = compare_labels def __getitem__(self, idx): - return self.sentence_indices1[idx], self.sentence_indices2[idx], self.labels[idx] + cmp_lbl = None if self.compare_labels is None else self.compare_labels[idx] + return self.sentence_indices1[idx], self.sentence_indices2[idx], self.labels[idx], cmp_lbl def __len__(self): return len(self.labels) @@ -53,10 +57,18 @@ def load_sick(): filename = os.path.join(config.sick_data, dataset, name) with open(filename) as f: for line in f: - indices = [embed_ids.get(word, padding_idx) for word in line.strip().split()] + indices = [embed_ids.get(word, -1) for word in line.strip().split()] + indices = list(filter(lambda x: x >= 0, indices)) sentence_indices.append(indices) return sentence_indices + def read_labels(filename): + labels = [] + with open(filename) as f: + for line in f: + labels.append([float(val) for val in line.split()]) + return labels + sets = [] embeddings = [] embed_ids = {} @@ -70,14 +82,13 @@ def load_sick(): embeddings.append([0.0] * 300) for dataset in ("train", "dev", "test"): - filename = os.path.join(config.sick_data, dataset, "sim_sparse.txt") - labels = [] - with open(filename) as f: - for line in f: - labels.append([float(val) for val in line.split()]) + sparse_filename = os.path.join(config.sick_data, dataset, "sim_sparse.txt") + truth_filename = os.path.join(config.sick_data, dataset, "sim.txt") + sparse_labels = read_labels(sparse_filename) + cmp_labels = read_labels(truth_filename) indices1 = fetch_indices("a.toks") indices2 = fetch_indices("b.toks") - sets.append(LabeledEmbeddedDataset(indices1, indices2, labels)) + sets.append(LabeledEmbeddedDataset(indices1, indices2, sparse_labels, cmp_labels)) embedding = nn.Embedding(len(embeddings), 300) embedding.weight.data.copy_(torch.Tensor(embeddings)) embedding.weight.requires_grad = False diff --git a/vdpwi/model.py b/vdpwi/model.py index 6b2f1ba..9d54e9b 100644 --- a/vdpwi/model.py +++ b/vdpwi/model.py @@ -17,7 +17,7 @@ class SerializableModule(nn.Module): class VDPWIConvNet(SerializableModule): def __init__(self, n_labels): super().__init__() - self.conv1 = nn.Conv2d(13, 128, 3, padding=1) + self.conv1 = nn.Conv2d(12, 128, 3, padding=1) self.conv2 = nn.Conv2d(128, 164, 3, padding=1) self.conv3 = nn.Conv2d(164, 192, 3, padding=1) self.conv4 = nn.Conv2d(192, 192, 3, padding=1) @@ -25,20 +25,16 @@ class VDPWIConvNet(SerializableModule): self.maxpool2 = nn.MaxPool2d(2, ceil_mode=True) self.dnn = nn.Linear(128, 128) self.output = nn.Linear(128, n_labels) + self.input_len = 32 def forward(self, x): - def pad_side(idx, max_size): - if max_size <= 32: - pad_len = 32 - x.size(idx) - elif max_size <= 48: - pad_len = 48 - x.size(idx) - else: - pad_len = 0 + def pad_side(idx): + pad_len = max(32 - x.size(idx), 0) return [0, pad_len] - padding = pad_side(3, max(x.size()[2:])) - padding.extend(pad_side(2, max(x.size()[2:]))) + padding = pad_side(3) + padding.extend(pad_side(2)) x = F.pad(x, padding) - x = x[:, :, :48, :48] + x = x[:, :, :32, :32] pool_final = nn.MaxPool2d(2, ceil_mode=True) if x.size(2) == 32 else nn.MaxPool2d(3, 1, ceil_mode=True) x = self.maxpool2(F.relu(self.conv1(x))) @@ -58,7 +54,7 @@ class VDPWIModel(SerializableModule): self.use_cuda = not config.cpu self.classifier_net = VDPWIConvNet(config.n_labels) if classifier_net is None else classifier_net - def compute_sim_cube(self, seq1, seq2): + def compute_sim_cube(self, seq1, seq2, truncate=None): def compute_sim(prism1, prism2): prism1_len = prism1.norm(dim=2) prism2_len = prism2.norm(dim=2) @@ -76,8 +72,7 @@ class VDPWIModel(SerializableModule): prism2 = prism2.permute(0, 1, 2).contiguous() return compute_sim(prism1, prism2) - sim_cube = Variable(torch.Tensor(13, seq1.size(0), seq2.size(0))) - sim_cube[12] = 0 + sim_cube = Variable(torch.Tensor(12, seq1.size(0), seq2.size(0))) if self.use_cuda: sim_cube = sim_cube.cuda() seq1_f = seq1[:, :self.hidden_dim] @@ -88,6 +83,8 @@ class VDPWIModel(SerializableModule): sim_cube[3:6] = compute_prism(seq1_f, seq2_f) sim_cube[6:9] = compute_prism(seq1_b, seq2_b) sim_cube[9:12] = compute_prism(seq1_f + seq1_b, seq2_f + seq2_b) + if truncate is not None: + sim_cube = sim_cube[:, :truncate, :truncate].contiguous() return sim_cube def compute_focus_cube(self, sim_cube): @@ -108,7 +105,6 @@ class VDPWIModel(SerializableModule): mask[:, int(pos1), int(pos2)] = 1 build_mask(9) build_mask(10) - mask[12, :, :] = 1 return mask * sim_cube def forward(self, x1, x2): @@ -122,7 +118,7 @@ class VDPWIModel(SerializableModule): seq2 = torch.cat([seq2f, seq2b], 2) seq1 = seq1.squeeze(0) # batch size assumed to be 1 seq2 = seq2.squeeze(0) - sim_cube = self.compute_sim_cube(seq1, seq2) + sim_cube = self.compute_sim_cube(seq1, seq2, truncate=self.classifier_net.input_len) focus_cube = self.compute_focus_cube(sim_cube) logits = self.classifier_net(focus_cube.unsqueeze(0)) return logits diff --git a/vdpwi/utils/__init__.py b/vdpwi/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vdpwi/utils/log.py b/vdpwi/utils/log.py new file mode 100644 index 0000000..82f5eef --- /dev/null +++ b/vdpwi/utils/log.py @@ -0,0 +1,26 @@ +import datetime +import sys + +from tensorboardX import SummaryWriter + +class LogWriter(object): + def __init__(self, run_name_fmt="run_{}"): + self.writer = SummaryWriter() + self.run_name = run_name_fmt.format(datetime.datetime.now().strftime("%Y-%m-%d-%H:%M:%S")) + self.train_idx = 0 + self.dev_idx = 0 + + def log_hyperparams(self): + self.writer.add_text("{}/hyperparams".format(self.run_name), " ".join(sys.argv)) + + def log_train_loss(self, loss): + self.writer.add_scalar("{}/train_loss".format(self.run_name), loss, self.train_idx) + self.train_idx += 1 + + def log_dev_metrics(self, pearsonr, spearmanr): + results = dict(pearsonr=pearsonr, spearmanr=spearmanr) + self.writer.add_scalars("{}/dev_metrics".format(self.run_name), results, self.dev_idx) + self.dev_idx += 1 + + def next(self): + self.i += 1