mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Update path name and README for NCE-SM (#95)
* update nce-sm * refactor code, update torchtext * use shared evaluation * refactor code, use shared data loader * refactor code * refactor code * refactor code according to Michael's great suggestions * update readme and requirement * update datasets and readme * update data loader * add space between + * update path name, update readme * update data loader and dataset name * refactor code * update readme
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
## NCE-SM model
|
||||
## NCE-SM-CNN model PyTorch Implementation
|
||||
|
||||
#### References:
|
||||
+ Aliaksei _S_everyn and Alessandro _M_oschitti. 2015. Learning to Rank Short Text Pairs with Convolutional Deep Neural
|
||||
@@ -15,8 +15,8 @@ cd text
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
Download the word2vec model from [here] (https://drive.google.com/file/d/0B2u_nClt6NbzUmhOZU55eEo4QWM/view?usp=sharing)
|
||||
and copy it to the `Castor/data/word2vec` folder.
|
||||
Download the word2vec model from [here] (https://drive.google.com/file/d/0B2u_nClt6NbzUmhOZU55eEo4QWM/view?usp=sharing)
|
||||
and copy it to the `Castor/data/word2vec` folder.
|
||||
|
||||
### Training the model
|
||||
|
||||
@@ -28,7 +28,7 @@ You can train the SM model for the 4 following configurations:
|
||||
|
||||
|
||||
```bash
|
||||
python train.py --no_cuda --mode rand --batch_size 64 --neg_num 8 --dev_every 50 --patience 1000
|
||||
python train.py --no_cuda --mode rand --batch_size 64 --neg_num 8 --dev_every 50 --patience 100 --dataset trec
|
||||
```
|
||||
|
||||
NB: pass `--no_cuda` to use CPU
|
||||
@@ -41,7 +41,7 @@ saves/static_best_model.pt
|
||||
### Testing the model
|
||||
|
||||
```
|
||||
python main.py --trained_model saves/TREC/multichannel_best_model.pt --batch_size 64 --no_cuda
|
||||
python main.py --trained_model saves/trec/multichannel_best_model.pt --batch_size 64 --no_cuda --dataset trec
|
||||
```
|
||||
|
||||
### Evaluation
|
||||
@@ -55,19 +55,56 @@ Metric |rand |static|non-static|multichannel
|
||||
MAP |0.7441 |0.7524|0.7688 |0.7641
|
||||
MRR |0.8172 |0.8012|0.8144 |0.8174
|
||||
|
||||
##### Max Neg Sample
|
||||
##### Pairwise + Random Sample with neg_num = 8
|
||||
|
||||
To be added
|
||||
Metric |rand |static|non-static|multichannel
|
||||
-------|-------|------|----------|------------
|
||||
MAP |0.7427 |0.7546|0.7614 | 0.7645
|
||||
MRR |0.8151 |0.8061|0.8162 | 0.8270
|
||||
|
||||
##### Pairwise + Max Neg Sample with neg_num = 8
|
||||
|
||||
Metric |rand |static|non-static|multichannel
|
||||
-------|-------|------|----------|------------
|
||||
MAP |0.7427 |0.7546|0.7716 |0.7794
|
||||
MRR |0.8151 |0.8061|0.8347 |0.8467
|
||||
MAP |0.7437 |0.7602|0.7752 |0.7664
|
||||
MRR |0.8151 |0.8109 |0.8270 |0.8347
|
||||
|
||||
|
||||
#### The performance on WikiQA dataset:
|
||||
|
||||
To be added
|
||||
##### Without NCE
|
||||
|
||||
Metric |rand |static|non-static|multichannel
|
||||
-------|-------|------|----------|------------
|
||||
MAP |0.6472 |0.6500 | 0.6620|0.6542
|
||||
MRR |0.664 |0.6693 | 0.6806|0.6722
|
||||
|
||||
##### Pairwise + Random Sample with neg_num = 8
|
||||
|
||||
Metric |rand |static|non-static|multichannel
|
||||
-------|-------|------|----------|------------
|
||||
MAP |0.6655 |0.6816|0.6697 |0.6739
|
||||
MRR |0.6831 |0.6992 |0.6929 |0.6925
|
||||
|
||||
##### Pairwise + Max Neg Sample with neg_num = 8
|
||||
|
||||
Metric |rand |static|non-static|multichannel
|
||||
-------|-------|------|----------|------------
|
||||
MAP |0.6687 |0.6796|0.6854 |0.6851
|
||||
MRR |0.6864 |0.6977 |0.7012 |0.7035
|
||||
|
||||
|
||||
## Optional Dependencies
|
||||
|
||||
To optionally visualize the learning curve during training, we make use of https://github.com/lanpa/tensorboard-pytorch to connect to [TensorBoard](https://github.com/tensorflow/tensorboard). These projects require TensorFlow as a dependency, so you need to install TensorFlow before running the commands below. After these are installed, just add `--tensorboard` when running `train.py` and open TensorBoard in the browser.
|
||||
|
||||
```sh
|
||||
pip install tensorboardX
|
||||
pip install tensorflow-tensorboard
|
||||
```
|
||||
|
||||
Usage:
|
||||
|
||||
```sh
|
||||
tensorboard --host 0.0.0.0 --port 5001 --logdir runs
|
||||
```
|
||||
@@ -3,13 +3,13 @@ from argparse import ArgumentParser
|
||||
def get_args():
|
||||
parser = ArgumentParser(description="SM CNN")
|
||||
parser.add_argument('--no_cuda', action='store_false', help='do not use cuda', dest='cuda')
|
||||
parser.add_argument('--gpu', type=int, default=0) # Use -1 for CPU
|
||||
parser.add_argument('--epochs', type=int, default=30)
|
||||
parser.add_argument('--gpu', type=int, default=-1) # Use -1 for CPU
|
||||
parser.add_argument('--epochs', type=int, default=10)
|
||||
parser.add_argument('--batch_size', type=int, default=64)
|
||||
parser.add_argument('--mode', type=str, default='static')
|
||||
parser.add_argument('--lr', type=float, default=0.95)
|
||||
parser.add_argument('--seed', type=int, default=3435)
|
||||
parser.add_argument('--dataset', type=str, default='TREC')
|
||||
parser.add_argument('--dataset', type=str, default='trec')
|
||||
parser.add_argument('--resume_snapshot', type=str, default=None)
|
||||
parser.add_argument('--dev_every', type=int, default=100)
|
||||
parser.add_argument('--log_every', type=int, default=10)
|
||||
@@ -20,7 +20,7 @@ def get_args():
|
||||
parser.add_argument('--words_dim', type=int, default=50)
|
||||
parser.add_argument('--dropout', type=float, default=0.5)
|
||||
parser.add_argument('--epoch_decay', type=int, default=15)
|
||||
parser.add_argument('--wordvec_dir', type=str, default='../../data/word2vec/')
|
||||
parser.add_argument('--wordvec_dir', type=str, default='../../../data/word2vec/')
|
||||
parser.add_argument('--vector_cache', type=str, default='word2vec.trecqa.pt')
|
||||
parser.add_argument('--trained_model', type=str, default="")
|
||||
parser.add_argument('--weight_decay',type=float, default=1e-5)
|
||||
@@ -29,6 +29,8 @@ def get_args():
|
||||
parser.add_argument('--neg_sample', type=str, default="random")
|
||||
parser.add_argument('--eps', type=float, default=1e-6)
|
||||
parser.add_argument('--optimizer', type=str, default="adadelta")
|
||||
parser.add_argument('--tensorboard', action='store_true', default=False,
|
||||
help='use TensorBoard to visualize training (default: false)')
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
@@ -7,9 +7,9 @@ import torch
|
||||
from torchtext import data
|
||||
|
||||
from args import get_args
|
||||
from trec_dataset import TrecDataset
|
||||
from utils.relevancy_metrics import get_map_mrr
|
||||
from datasets.trecqa import TRECQA
|
||||
from datasets.wikiqa import WikiQA
|
||||
from train import UnknownWordVecCache
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -36,22 +36,27 @@ if torch.cuda.is_available() and args.cuda:
|
||||
if torch.cuda.is_available() and not args.cuda:
|
||||
logger.info("Warning: You have Cuda but do not use it. You are using CPU for training")
|
||||
|
||||
|
||||
if config.dataset == 'TREC':
|
||||
dataset_root = os.path.join(os.pardir, 'data', 'TrecQA/')
|
||||
train_iter, dev_iter, test_iter = TRECQA.iters(dataset_root, args.vector_cache, args.wordvec_dir, batch_size=args.batch_size, pt_file=True, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
if args.dataset == "trec":
|
||||
dataset_cls = TRECQA
|
||||
dataset_root = os.path.join(os.pardir, os.pardir, os.pardir, 'data', 'TrecQA/')
|
||||
elif args.dataset == "wiki":
|
||||
dataset_cls = WikiQA
|
||||
dataset_root = os.path.join(os.pardir, os.pardir, os.pardir, 'data', 'WikiQA/')
|
||||
else:
|
||||
logger.info("Unsupported dataset")
|
||||
exit()
|
||||
|
||||
train_iter, dev_iter, test_iter = dataset_cls.iters(dataset_root, args.vector_cache, args.wordvec_dir, batch_size=args.batch_size, pt_file=True, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
|
||||
config.target_class = 2
|
||||
config.questions_num = len(TRECQA.TEXT_FIELD.vocab)
|
||||
config.answers_num = len(TRECQA.TEXT_FIELD.vocab)
|
||||
config.questions_num = len(dataset_cls.TEXT_FIELD.vocab)
|
||||
config.answers_num = len(dataset_cls.TEXT_FIELD.vocab)
|
||||
|
||||
if args.cuda:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage,location: storage)
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage)
|
||||
|
||||
|
||||
|
||||
def predict(test_mode, dataset_iter):
|
||||
@@ -74,6 +79,7 @@ def predict(test_mode, dataset_iter):
|
||||
|
||||
logger.info("{} {}".format(dev_map, dev_mrr))
|
||||
|
||||
|
||||
# Run the model on the dev set
|
||||
predict('dev', dataset_iter=dev_iter)
|
||||
|
||||
@@ -5,13 +5,16 @@ import random
|
||||
import heapq
|
||||
import operator
|
||||
import logging
|
||||
import pprint
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchtext import data
|
||||
from torch.nn import functional as F
|
||||
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
||||
|
||||
from datasets.trecqa import TRECQA
|
||||
from datasets.wikiqa import WikiQA
|
||||
from args import get_args
|
||||
from model import SmPlusPlus, PairwiseConv
|
||||
from utils.relevancy_metrics import get_map_mrr
|
||||
@@ -29,7 +32,8 @@ class UnknownWordVecCache(object):
|
||||
size_tup = tuple(tensor.size())
|
||||
if size_tup not in cls.cache:
|
||||
cls.cache[size_tup] = torch.Tensor(tensor.size())
|
||||
cls.cache[size_tup].uniform_(-0.05, 0.05)
|
||||
# cls.cache[size_tup].uniform_(-0.05, 0.05)
|
||||
cls.cache[size_tup].uniform_(-0.25, 0.25)
|
||||
return cls.cache[size_tup]
|
||||
|
||||
|
||||
@@ -47,8 +51,12 @@ def train_sm():
|
||||
config = args
|
||||
torch.backends.cudnn.deterministic = True
|
||||
|
||||
logger.info(pprint.pformat(vars(args)))
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(args.seed)
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
if not args.cuda:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
@@ -57,22 +65,28 @@ def train_sm():
|
||||
torch.cuda.manual_seed(args.seed)
|
||||
if torch.cuda.is_available() and not args.cuda:
|
||||
logger.info("You have Cuda but you're using CPU for training.")
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
|
||||
dataset_root = os.path.join(os.pardir, 'data', 'TrecQA/')
|
||||
train_iter, dev_iter, test_iter = TRECQA.iters(dataset_root, args.vector_cache, args.wordvec_dir, batch_size=args.batch_size,
|
||||
pt_file=True, device=args.gpu, unk_init=UnknownWordVecCache.unk) #
|
||||
if args.dataset == "trec":
|
||||
dataset_cls = TRECQA
|
||||
dataset_root = os.path.join(os.pardir, os.pardir, os.pardir, 'data', 'TrecQA/')
|
||||
elif args.dataset == "wiki":
|
||||
dataset_cls = WikiQA
|
||||
dataset_root = os.path.join(os.pardir, os.pardir, os.pardir, 'data', 'WikiQA/')
|
||||
|
||||
index2text = np.array(TRECQA.TEXT_FIELD.vocab.itos)
|
||||
train_iter, dev_iter, test_iter = dataset_cls.iters(dataset_root, args.vector_cache, args.wordvec_dir,
|
||||
batch_size=args.batch_size,
|
||||
pt_file=True, device=args.gpu,
|
||||
unk_init=UnknownWordVecCache.unk) #
|
||||
|
||||
index2text = np.array(dataset_cls.TEXT_FIELD.vocab.itos)
|
||||
|
||||
config.target_class = 2
|
||||
config.questions_num = TRECQA.VOCAB_SIZE
|
||||
config.answers_num = TRECQA.VOCAB_SIZE
|
||||
config.questions_num = dataset_cls.VOCAB_SIZE
|
||||
config.answers_num = dataset_cls.VOCAB_SIZE
|
||||
|
||||
logger.info("index2text: {}".format(index2text))
|
||||
logger.info("Dataset: {}, Mode: {}".format(args.dataset, args.mode))
|
||||
logger.info("VOCAB num: {}".format(TRECQA.VOCAB_SIZE))
|
||||
logger.info("VOCAB num: {}".format(dataset_cls.VOCAB_SIZE))
|
||||
|
||||
if args.resume_snapshot:
|
||||
if args.cuda:
|
||||
@@ -81,10 +95,10 @@ def train_sm():
|
||||
pw_model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage)
|
||||
else:
|
||||
model = SmPlusPlus(config)
|
||||
model.static_question_embed.weight.data.copy_(TRECQA.TEXT_FIELD.vocab.vectors)
|
||||
model.nonstatic_question_embed.weight.data.copy_(TRECQA.TEXT_FIELD.vocab.vectors)
|
||||
model.static_answer_embed.weight.data.copy_(TRECQA.TEXT_FIELD.vocab.vectors)
|
||||
model.nonstatic_answer_embed.weight.data.copy_(TRECQA.TEXT_FIELD.vocab.vectors)
|
||||
model.static_question_embed.weight.data.copy_(dataset_cls.TEXT_FIELD.vocab.vectors)
|
||||
model.nonstatic_question_embed.weight.data.copy_(dataset_cls.TEXT_FIELD.vocab.vectors)
|
||||
model.static_answer_embed.weight.data.copy_(dataset_cls.TEXT_FIELD.vocab.vectors)
|
||||
model.nonstatic_answer_embed.weight.data.copy_(dataset_cls.TEXT_FIELD.vocab.vectors)
|
||||
|
||||
if args.cuda:
|
||||
model.cuda()
|
||||
@@ -124,11 +138,22 @@ def train_sm():
|
||||
log_template = ' '.join('{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>11.6f},{:>11.6f},'.split(','))
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.save_path, args.dataset), exist_ok=True)
|
||||
print(header)
|
||||
logger.info(header)
|
||||
|
||||
filename = "grid_{dataset}_lr_{learning_rate}_eps_{eps}_reg_{reg}_mode_{mode}_device_{dev}.txt".format(
|
||||
learning_rate=args.lr, eps=args.eps, reg=args.weight_decay, dev=args.gpu, dataset=args.dataset, mode=args.mode)
|
||||
|
||||
if args.tensorboard:
|
||||
from tensorboardX import SummaryWriter
|
||||
writer = SummaryWriter(log_dir=None, comment=filename)
|
||||
|
||||
dev_index = 0
|
||||
train_index = 0
|
||||
|
||||
# scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.3, patience=8)
|
||||
while True:
|
||||
if early_stop:
|
||||
logger.log("Early Stopping. Epoch: {}, Best Dev Loss: {}".format(epoch, best_dev_loss))
|
||||
logger.info("Early Stopping. Epoch: {}, Best Dev Map: {}, Best Dev Mrr: {}".format(epoch, best_dev_map, best_dev_mrr))
|
||||
break
|
||||
epoch += 1
|
||||
train_iter.init_epoch()
|
||||
@@ -229,6 +254,7 @@ def train_sm():
|
||||
|
||||
# pack the selected pos and neg samples into the torchtext batch and train
|
||||
if epoch != 1:
|
||||
train_index += 1
|
||||
true_batch_size = len(new_train_neg["answer"])
|
||||
if true_batch_size != 0:
|
||||
for j in range(true_batch_size):
|
||||
@@ -266,6 +292,7 @@ def train_sm():
|
||||
qids = []
|
||||
predictions = []
|
||||
labels = []
|
||||
dev_index += 1
|
||||
|
||||
for dev_batch_idx, dev_batch in enumerate(dev_iter):
|
||||
'''
|
||||
@@ -273,7 +300,8 @@ def train_sm():
|
||||
but dev singlely is equal to dev_size = 1
|
||||
'''
|
||||
scores = pw_model.convModel(dev_batch)
|
||||
scores = pw_model.linearLayer(scores)
|
||||
# scores = pw_model.linearLayer(scores)
|
||||
scores = pw_model.predict(scores)
|
||||
qid_array = np.transpose(dev_batch.id.cpu().data.numpy())
|
||||
score_array = scores.cpu().data.numpy().reshape(-1)
|
||||
true_label_array = np.transpose(dev_batch.label.cpu().data.numpy())
|
||||
@@ -283,14 +311,42 @@ def train_sm():
|
||||
labels.extend(true_label_array.tolist())
|
||||
|
||||
dev_map, dev_mrr = get_map_mrr(qids, predictions, labels)
|
||||
print(dev_log_template.format(time.time() - start,
|
||||
logger.info(dev_log_template.format(time.time() - start,
|
||||
epoch, iterations, 1 + batch_idx, len(train_iter),
|
||||
100. * (1 + batch_idx) / len(train_iter),
|
||||
loss_num, acc / tot, dev_map, dev_mrr))
|
||||
|
||||
qids = []
|
||||
predictions = []
|
||||
labels = []
|
||||
for test_batch_idx, test_batch in enumerate(test_iter):
|
||||
'''
|
||||
# dev singlely or in a batch? -> in a batch
|
||||
but dev singlely is equal to dev_size = 1
|
||||
'''
|
||||
scores = pw_model.convModel(test_batch)
|
||||
# scores = pw_model.linearLayer(scores)
|
||||
scores = pw_model.predict(scores)
|
||||
qid_array = np.transpose(test_batch.id.cpu().data.numpy())
|
||||
score_array = scores.cpu().data.numpy().reshape(-1)
|
||||
true_label_array = np.transpose(test_batch.label.cpu().data.numpy())
|
||||
|
||||
qids.extend(qid_array.tolist())
|
||||
predictions.extend(score_array.tolist())
|
||||
labels.extend(true_label_array.tolist())
|
||||
|
||||
if args.tensorboard:
|
||||
writer.add_scalar('{}/dev/map'.format(args.dataset), dev_map, dev_index)
|
||||
writer.add_scalar('{}/dev/mrr'.format(args.dataset), dev_mrr, dev_index)
|
||||
writer.add_scalar('{}/lr'.format(args.dataset),
|
||||
optimizer.param_groups[0]['lr'], dev_index)
|
||||
writer.add_scalar('{}/train/loss'.format(args.dataset), loss_num, dev_index)
|
||||
|
||||
if best_dev_mrr < dev_mrr:
|
||||
snapshot_path = os.path.join(args.save_path, args.dataset, args.mode + '_best_model.pt')
|
||||
torch.save(pw_model, snapshot_path)
|
||||
iters_not_improved = 0
|
||||
best_dev_map = dev_map
|
||||
best_dev_mrr = dev_mrr
|
||||
else:
|
||||
iters_not_improved += 1
|
||||
@@ -298,14 +354,19 @@ def train_sm():
|
||||
early_stop = True
|
||||
break
|
||||
|
||||
# scheduler.step(dev_mrr)
|
||||
if iterations % args.log_every == 1 and epoch != 1:
|
||||
# logger.info progress message
|
||||
print(log_template.format(time.time() - start,
|
||||
logger.info(log_template.format(time.time() - start,
|
||||
epoch, iterations, 1 + batch_idx, len(train_iter),
|
||||
100. * (1 + batch_idx) / len(train_iter),
|
||||
loss_num, acc / tot))
|
||||
|
||||
|
||||
acc = 0
|
||||
tot = 0
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
train_sm()
|
||||
train_sm()
|
||||
Reference in New Issue
Block a user