mirror of
https://github.com/wassname/Castor.git
synced 2026-09-10 11:40:44 +08:00
Kim CNN OOP Refactoring (#124)
* Nuke obsolete artifacts * Refactor Kim CNN * Make kim_cnn a module * Fix bugs * Update README * Add choices to dataset arg * update for sst2 update sst.py update sst.py * Add Kim CNN dataset choices to args.py * Update tuned SST-1 accuracy
This commit is contained in:
@@ -1,10 +1,12 @@
|
||||
from .evaluators.sick_evaluator import SICKEvaluator
|
||||
from .evaluators.msrvid_evaluator import MSRVIDEvaluator
|
||||
from .evaluators.sst_evaluator import SSTEvaluator
|
||||
from .evaluators.trecqa_evaluator import TRECQAEvaluator
|
||||
from .evaluators.wikiqa_evaluator import WikiQAEvaluator
|
||||
from nce.nce_pairwise_mp.evaluators.trecqa_evaluator import TRECQAEvaluatorNCE
|
||||
from nce.nce_pairwise_mp.evaluators.wikiqa_evaluator import WikiQAEvaluatorNCE
|
||||
|
||||
|
||||
class EvaluatorFactory(object):
|
||||
"""
|
||||
Get the corresponding Evaluator class for a particular dataset.
|
||||
@@ -12,6 +14,8 @@ class EvaluatorFactory(object):
|
||||
evaluator_map = {
|
||||
'sick': SICKEvaluator,
|
||||
'msrvid': MSRVIDEvaluator,
|
||||
'SST-1': SSTEvaluator,
|
||||
'SST-2': SSTEvaluator,
|
||||
'trecqa': TRECQAEvaluator,
|
||||
'wikiqa': WikiQAEvaluator
|
||||
}
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .evaluator import Evaluator
|
||||
|
||||
|
||||
class SSTEvaluator(Evaluator):
|
||||
|
||||
def get_scores(self):
|
||||
self.model.eval()
|
||||
self.data_loader.init_epoch()
|
||||
n_dev_correct = 0
|
||||
total_loss = 0
|
||||
|
||||
for batch_idx, batch in enumerate(self.data_loader):
|
||||
scores = self.model(batch)
|
||||
n_dev_correct += (
|
||||
torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum().item()
|
||||
total_loss += F.cross_entropy(scores, batch.label, size_average=False).item()
|
||||
|
||||
accuracy = 100. * n_dev_correct / len(self.data_loader.dataset.examples)
|
||||
avg_loss = total_loss / len(self.data_loader.dataset.examples)
|
||||
|
||||
return [accuracy, avg_loss], ['accuracy', 'cross_entropy_loss']
|
||||
@@ -2,6 +2,7 @@ from .trainers.sick_trainer import SICKTrainer
|
||||
from .trainers.msrvid_trainer import MSRVIDTrainer
|
||||
from .trainers.trecqa_trainer import TRECQATrainer
|
||||
from .trainers.wikiqa_trainer import WikiQATrainer
|
||||
from .trainers.sst_trainer import SSTTrainer
|
||||
from nce.nce_pairwise_mp.trainers.trecqa_trainer import TRECQATrainerNCE
|
||||
from nce.nce_pairwise_mp.trainers.wikiqa_trainer import WikiQATrainerNCE
|
||||
|
||||
@@ -13,6 +14,8 @@ class TrainerFactory(object):
|
||||
trainer_map = {
|
||||
'sick': SICKTrainer,
|
||||
'msrvid': MSRVIDTrainer,
|
||||
'SST-1': SSTTrainer,
|
||||
'SST-2': SSTTrainer,
|
||||
'trecqa': TRECQATrainer,
|
||||
'wikiqa': WikiQATrainer
|
||||
}
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
import time
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .trainer import Trainer
|
||||
from utils.serialization import save_checkpoint
|
||||
|
||||
|
||||
class SSTTrainer(Trainer):
|
||||
|
||||
def __init__(self, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator):
|
||||
super(SSTTrainer, self).__init__(model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator)
|
||||
self.early_stop = False
|
||||
self.best_dev_acc = 0
|
||||
self.iterations = 0
|
||||
self.iters_not_improved = 0
|
||||
self.start = None
|
||||
self.log_template = ' '.join(
|
||||
'{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{},{:12.4f},{}'.split(','))
|
||||
self.dev_log_template = ' '.join('{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{:8.6f},{:12.4f},{:12.4f}'.split(','))
|
||||
|
||||
def train_epoch(self, epoch):
|
||||
self.train_loader.init_epoch()
|
||||
n_correct, n_total = 0, 0
|
||||
for batch_idx, batch in enumerate(self.train_loader):
|
||||
self.iterations += 1
|
||||
self.model.train()
|
||||
self.optimizer.zero_grad()
|
||||
scores = self.model(batch)
|
||||
n_correct += (torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum().item()
|
||||
n_total += batch.batch_size
|
||||
train_acc = 100. * n_correct / n_total
|
||||
|
||||
loss = F.cross_entropy(scores, batch.label)
|
||||
loss.backward()
|
||||
|
||||
self.optimizer.step()
|
||||
|
||||
# Evaluate performance on validation set
|
||||
if self.iterations % self.dev_log_interval == 1:
|
||||
dev_acc, dev_loss = self.dev_evaluator.get_scores()[0]
|
||||
print(self.dev_log_template.format(time.time() - self.start,
|
||||
epoch, self.iterations, 1 + batch_idx, len(self.train_loader),
|
||||
100. * (1 + batch_idx) / len(self.train_loader), loss.item(),
|
||||
dev_loss, train_acc, dev_acc))
|
||||
|
||||
|
||||
# Update validation results
|
||||
if dev_acc > self.best_dev_acc:
|
||||
self.iters_not_improved = 0
|
||||
self.best_dev_acc = dev_acc
|
||||
snapshot_path = os.path.join(self.model_outfile, self.train_loader.dataset.NAME, self.model.mode + '_best_model.pt')
|
||||
torch.save(self.model, snapshot_path)
|
||||
else:
|
||||
self.iters_not_improved += 1
|
||||
if self.iters_not_improved >= self.patience:
|
||||
self.early_stop = True
|
||||
break
|
||||
|
||||
if self.iterations % self.log_interval == 1:
|
||||
# print progress message
|
||||
print(self.log_template.format(time.time() - self.start,
|
||||
epoch, self.iterations, 1 + batch_idx, len(self.train_loader),
|
||||
100. * (1 + batch_idx) / len(self.train_loader), loss.item(), ' ' * 8,
|
||||
train_acc, ' ' * 12))
|
||||
|
||||
def train(self, epochs):
|
||||
self.start = time.time()
|
||||
header = ' Time Epoch Iteration Progress (%Epoch) Loss Dev/Loss Accuracy Dev/Accuracy'
|
||||
# model_outfile is actually a directory, using model_outfile to conform to Trainer naming convention
|
||||
os.makedirs(self.model_outfile, exist_ok=True)
|
||||
os.makedirs(os.path.join(self.model_outfile, self.train_loader.dataset.NAME), exist_ok=True)
|
||||
print(header)
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
if self.early_stop:
|
||||
print("Early Stopping. Epoch: {}, Best Dev Acc: {}".format(epoch, self.best_dev_acc))
|
||||
break
|
||||
self.train_epoch(epoch)
|
||||
+13
-11
@@ -7,20 +7,21 @@ class Trainer(object):
|
||||
def __init__(self, model, embedding, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None):
|
||||
self.model = model
|
||||
self.embedding = embedding
|
||||
self.optimizer = trainer_config['optimizer']
|
||||
self.optimizer = trainer_config.get('optimizer')
|
||||
self.train_loader = train_loader
|
||||
self.batch_size = trainer_config['batch_size']
|
||||
self.log_interval = trainer_config['log_interval']
|
||||
self.model_outfile = trainer_config['model_outfile']
|
||||
self.lr_reduce_factor = trainer_config['lr_reduce_factor']
|
||||
self.patience = trainer_config['patience']
|
||||
self.use_tensorboard = trainer_config['tensorboard']
|
||||
self.batch_size = trainer_config.get('batch_size')
|
||||
self.log_interval = trainer_config.get('log_interval')
|
||||
self.dev_log_interval = trainer_config.get('dev_log_interval')
|
||||
self.model_outfile = trainer_config.get('model_outfile')
|
||||
self.lr_reduce_factor = trainer_config.get('lr_reduce_factor')
|
||||
self.patience = trainer_config.get('patience')
|
||||
self.use_tensorboard = trainer_config.get('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'])
|
||||
self.logger = trainer_config['logger']
|
||||
self.logger = trainer_config.get('logger')
|
||||
|
||||
self.train_evaluator = train_evaluator
|
||||
self.test_evaluator = test_evaluator
|
||||
@@ -28,9 +29,10 @@ class Trainer(object):
|
||||
|
||||
def evaluate(self, evaluator, dataset_name):
|
||||
scores, metric_names = evaluator.get_scores()
|
||||
self.logger.info('Evaluation metrics for {}:'.format(dataset_name))
|
||||
self.logger.info('\t'.join([' '] + metric_names))
|
||||
self.logger.info('\t'.join([dataset_name] + list(map(str, scores))))
|
||||
if self.logger is not None:
|
||||
self.logger.info('Evaluation metrics for {}:'.format(dataset_name))
|
||||
self.logger.info('\t'.join([' '] + metric_names))
|
||||
self.logger.info('\t'.join([dataset_name] + list(map(str, scores))))
|
||||
return scores
|
||||
|
||||
def get_sentence_embeddings(self, batch):
|
||||
|
||||
+44
-1
@@ -16,7 +16,7 @@ def clean_str_sst(string):
|
||||
|
||||
|
||||
class SST1(TabularDataset):
|
||||
NAME = 'sst-1'
|
||||
NAME = 'SST-1'
|
||||
NUM_CLASSES = 5
|
||||
|
||||
TEXT_FIELD = Field(batch_first=True, tokenize=clean_str_sst)
|
||||
@@ -55,3 +55,46 @@ class SST1(TabularDataset):
|
||||
|
||||
return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle,
|
||||
sort_within_batch=True, device=device)
|
||||
|
||||
class SST2(TabularDataset):
|
||||
NAME = 'SST-2'
|
||||
NUM_CLASSES = 5
|
||||
|
||||
TEXT_FIELD = Field(batch_first=True, tokenize=clean_str_sst)
|
||||
LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True)
|
||||
|
||||
@staticmethod
|
||||
def sort_key(ex):
|
||||
return len(ex.text)
|
||||
|
||||
@classmethod
|
||||
def splits(cls, path, train='stsa.binary.phrases.train', validation='stsa.binary.dev', test='stsa.binary.test', **kwargs):
|
||||
return super(SST2, cls).splits(
|
||||
path, train=train, validation=validation, test=test,
|
||||
format='tsv', fields=[('label', cls.LABEL_FIELD), ('text', cls.TEXT_FIELD)]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def iters(cls, path, vectors_name, vectors_cache, batch_size=64, shuffle=True, device=0, vectors=None,
|
||||
unk_init=torch.Tensor.zero_):
|
||||
"""
|
||||
:param path: directory containing train, test, dev files
|
||||
:param vectors_name: name of word vectors file
|
||||
:param vectors_cache: path to directory containing word vectors file
|
||||
:param batch_size: batch size
|
||||
:param device: GPU device
|
||||
:param vectors: custom vectors - either predefined torchtext vectors or your own custom Vector classes
|
||||
:param unk_init: function used to generate vector for OOV words
|
||||
:return:
|
||||
"""
|
||||
if vectors is None:
|
||||
vectors = Vectors(name=vectors_name, cache=vectors_cache, unk_init=unk_init)
|
||||
|
||||
train, val, test = cls.splits(path)
|
||||
|
||||
cls.TEXT_FIELD.build_vocab(train, val, test, min_freq=2, vectors=vectors)
|
||||
|
||||
return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle,
|
||||
sort_within_batch=True, device=device)
|
||||
|
||||
|
||||
|
||||
+38
-40
@@ -10,56 +10,35 @@ Implementation for Convolutional Neural Networks for Sentence Classification of
|
||||
- multichannel: A model with two sets of word vectors. Each set of vectors is treated as a 'channel' and each filter is applied to both channels, but gradients are back-propagated only through one of the channels. Hence the model is able to fine-tune one set of vectors while keeping the other static. Both channels are initialized with word2vec.# text-classification-cnn
|
||||
Implementation for Convolutional Neural Networks for Sentence Classification of [Kim (2014)](https://arxiv.org/abs/1408.5882) with PyTorch.
|
||||
|
||||
## Requirement
|
||||
|
||||
Assuming you already have PyTorch, just install torchtext (`pip install torchtext==0.2.0`)
|
||||
|
||||
## Quick Start
|
||||
|
||||
To get the dataset, you can run this.
|
||||
```
|
||||
cd kim_cnn
|
||||
bash get_data.sh
|
||||
```
|
||||
|
||||
To run the model on SST-1 dataset on multichannel, just run the following code.
|
||||
To run the model on SST-1 dataset on multichannel, just run the following from the Castor working directory.
|
||||
|
||||
```
|
||||
python train.py --mode multichannel
|
||||
python -m kim_cnn --mode multichannel
|
||||
```
|
||||
|
||||
The file will be saved in
|
||||
The file will be saved in
|
||||
|
||||
```
|
||||
saves/best_model.pt
|
||||
kim_cnn/saves/best_model.pt
|
||||
```
|
||||
|
||||
To test the model, you can use the following command.
|
||||
|
||||
```
|
||||
python main.py --trained_model saves/best_model.pt --mode multichannel
|
||||
python -m kim_cnn --trained_model kim_cnn/saves/SST-1/multichannel_best_model.pt --mode multichannel
|
||||
```
|
||||
|
||||
## Dataset
|
||||
|
||||
|
||||
## Dataset and Embeddings
|
||||
|
||||
We experiment the model on the following three datasets.
|
||||
We experiment the model on the following datasets.
|
||||
|
||||
- SST-1: Keep the original splits and train with phrase level dataset and test on sentence level dataset.
|
||||
|
||||
**word2vec.sst-1.pt** is a subset of word2vector. We just select the word appearing in the SST-1 dataset and generate this file with the **vector_preprocess.py**(you will get this after you run get_data.sh or you can download [here](https://raw.githubusercontent.com/Impavidity/kim_cnn/master/vector_preprocess.py)) You can select these from any kind of word embedding text file and generate in following format.
|
||||
```
|
||||
word vector_in_one_line
|
||||
```
|
||||
and then run
|
||||
```
|
||||
python vector_preprocess.py file_in embed.pt
|
||||
```
|
||||
Here you can get *embed.pt* for the embedding file. Remember change the argument in *args.py* file with your own embedding.
|
||||
|
||||
## Settings
|
||||
Adadelta is used for training.
|
||||
|
||||
Adadelta is used for training.
|
||||
|
||||
## Training Time
|
||||
|
||||
@@ -78,21 +57,40 @@ torch.backends.cudnn.enabled = False
|
||||
```
|
||||
but this will take ~6-7x training time.
|
||||
|
||||
## Results
|
||||
## SST-1 Dataset Results
|
||||
|
||||
Deterministic Algorithm for CNN.
|
||||
**Random**
|
||||
|
||||
| Dev Accuracy on SST-1 | rand | static | non-static | multichannel |
|
||||
|:--------------------------:|:-----------:|:-----------:|:-------------:|:---------------:|
|
||||
| My-Implementation | 42.597639| 48.773842| 48.864668 | 49.046322 |
|
||||
```
|
||||
python -m kim_cnn --mode rand --lr 0.8337 --weight_decay 0.0008987 --dropout 0.4
|
||||
```
|
||||
|
||||
| Test Accuracy on SST-1| rand | static | non-static | multichannel |
|
||||
|:--------------------------:|:-----------:|:-----------:|:-------------:|:---------------:|
|
||||
| Kim-Implementation | 45.0 | 45.5 | 48.0 | 47.4 |
|
||||
| My- Implementation | 39.683258 | 45.972851| 48.914027| 47.330317 |
|
||||
**Static**
|
||||
|
||||
```
|
||||
python -m kim_cnn --mode static --lr 0.8641 --weight_decay 1.44e-05 --dropout 0.3
|
||||
```
|
||||
|
||||
**Non-static**
|
||||
|
||||
```
|
||||
python -m kim_cnn --mode non-static --lr 0.371 --weight_decay 1.84e-05 --dropout 0.4
|
||||
```
|
||||
|
||||
**Multichannel**
|
||||
|
||||
```
|
||||
python -m kim_cnn --mode multichannel --lr 0.2532 --weight_decay 3.95e-05 --dropout 0.1
|
||||
```
|
||||
|
||||
Using deterministic algorithm for cuDNN.
|
||||
|
||||
| Test Accuracy on SST-1 | rand | static | non-static | multichannel |
|
||||
|:------------------------------:|:----------:|:------------:|:--------------:|:---------------:|
|
||||
| Paper | 45.0 | 45.5 | 48.0 | 47.4 |
|
||||
| PyTorch using above configs | 41.5 | 44.7 | 47.4 | 47.5 |
|
||||
|
||||
## TODO
|
||||
|
||||
- More experiments on SST-2 and subjectivity
|
||||
- Parameters tuning
|
||||
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
from torchtext import data
|
||||
import os
|
||||
|
||||
|
||||
class SST1Dataset(data.TabularDataset):
|
||||
dirname = 'data'
|
||||
@classmethod
|
||||
def splits(cls, text_field, label_field,
|
||||
train='phrases.train.tsv', validation='dev.tsv', test='test.tsv'):
|
||||
prefix_name = 'stsa.fine.'
|
||||
path = './data'
|
||||
return super(SST1Dataset, cls).splits(
|
||||
path, train=prefix_name + train, validation=prefix_name + validation, test=prefix_name + test,
|
||||
format='TSV', fields=[('label', label_field), ('text', text_field)]
|
||||
)
|
||||
@@ -0,0 +1,129 @@
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from common.evaluation import EvaluatorFactory
|
||||
from common.train import TrainerFactory
|
||||
from datasets.sst import SST1
|
||||
from datasets.sst import SST2
|
||||
from kim_cnn.args import get_args
|
||||
from kim_cnn.model import KimCNN
|
||||
|
||||
|
||||
def get_logger():
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
ch = logging.StreamHandler()
|
||||
ch.setLevel(logging.DEBUG)
|
||||
formatter = logging.Formatter('%(levelname)s - %(message)s')
|
||||
ch.setFormatter(formatter)
|
||||
logger.addHandler(ch)
|
||||
|
||||
return logger
|
||||
|
||||
|
||||
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))))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# Set default configuration in : args.py
|
||||
args = get_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(args.seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
if not args.cuda:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
print("Note: You are using GPU for training")
|
||||
torch.cuda.set_device(args.gpu)
|
||||
torch.cuda.manual_seed(args.seed)
|
||||
if torch.cuda.is_available() and not args.cuda:
|
||||
print("Warning: You have Cuda but not use it. You are using CPU for training.")
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
logger = get_logger()
|
||||
|
||||
# Set up the data for training SST-1
|
||||
if args.dataset == 'SST-1':
|
||||
train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu)
|
||||
# Set up the data for training SST-2
|
||||
elif args.dataset == 'SST-2':
|
||||
train_iter, dev_iter, test_iter = SST2.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
config = deepcopy(args)
|
||||
config.dataset = train_iter.dataset
|
||||
config.target_class = train_iter.dataset.NUM_CLASSES
|
||||
config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab)
|
||||
|
||||
print("Dataset {} Mode {}".format(args.dataset, args.mode))
|
||||
print("VOCAB num",len(train_iter.dataset.TEXT_FIELD.vocab))
|
||||
print("LABEL.target_class:", train_iter.dataset.NUM_CLASSES)
|
||||
print("Train instance", len(train_iter.dataset))
|
||||
print("Dev instance", len(dev_iter.dataset))
|
||||
print("Test instance", len(test_iter.dataset))
|
||||
|
||||
if args.resume_snapshot:
|
||||
if args.cuda:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage)
|
||||
else:
|
||||
model = KimCNN(config)
|
||||
if args.cuda:
|
||||
model.cuda()
|
||||
print("Shift model to GPU")
|
||||
|
||||
parameter = filter(lambda p: p.requires_grad, model.parameters())
|
||||
optimizer = torch.optim.Adadelta(parameter, lr=args.lr, weight_decay=args.weight_decay)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
trainer_config = {
|
||||
'optimizer': optimizer,
|
||||
'batch_size': args.batch_size,
|
||||
'log_interval': args.log_every,
|
||||
'dev_log_interval': args.dev_every,
|
||||
'patience': args.patience,
|
||||
'model_outfile': args.save_path, # actually a directory, using model_outfile to conform to Trainer naming convention
|
||||
'logger': logger
|
||||
}
|
||||
trainer = TrainerFactory.get_trainer(args.dataset, model, None, train_iter, trainer_config, train_evaluator, test_evaluator, dev_evaluator)
|
||||
|
||||
if not args.trained_model:
|
||||
trainer.train(args.epochs)
|
||||
else:
|
||||
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)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
evaluate_dataset('dev', SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
evaluate_dataset('dev', SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
+7
-6
@@ -2,30 +2,31 @@ import os
|
||||
|
||||
from argparse import ArgumentParser
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = ArgumentParser(description="Kim 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('--batch_size', type=int, default=1000)
|
||||
parser.add_argument('--mode', type=str, default='multichannel')
|
||||
parser.add_argument('--batch_size', type=int, default=1024)
|
||||
parser.add_argument('--mode', type=str, default='multichannel', choices=['rand', 'static', 'non-static', 'multichannel'])
|
||||
parser.add_argument('--lr', type=float, default=1.0)
|
||||
parser.add_argument('--seed', type=int, default=3435)
|
||||
parser.add_argument('--dataset', type=str, default='SST-1')
|
||||
parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2'])
|
||||
parser.add_argument('--resume_snapshot', type=str, default=None)
|
||||
parser.add_argument('--dev_every', type=int, default=30)
|
||||
parser.add_argument('--log_every', type=int, default=10)
|
||||
parser.add_argument('--patience', type=int, default=50)
|
||||
parser.add_argument('--save_path', type=str, default='saves')
|
||||
parser.add_argument('--save_path', type=str, default='kim_cnn/saves')
|
||||
parser.add_argument('--output_channel', type=int, default=100)
|
||||
parser.add_argument('--words_dim', type=int, default=300)
|
||||
parser.add_argument('--embed_dim', type=int, default=300)
|
||||
parser.add_argument('--dropout', type=float, default=0.5)
|
||||
parser.add_argument('--epoch_decay', type=int, default=15)
|
||||
parser.add_argument('--data_dir', help='word vectors directory',
|
||||
default=os.path.join(os.pardir, os.pardir, 'Castor-data', 'datasets', 'SST'))
|
||||
default=os.path.join(os.pardir, 'Castor-data', 'datasets', 'SST'))
|
||||
parser.add_argument('--word_vectors_dir', help='word vectors directory',
|
||||
default=os.path.join(os.pardir, os.pardir, 'Castor-data', 'embeddings', 'word2vec'))
|
||||
default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'word2vec'))
|
||||
parser.add_argument('--word_vectors_file', help='word vectors filename', default='GoogleNews-vectors-negative300.txt')
|
||||
parser.add_argument('--trained_model', type=str, default="")
|
||||
parser.add_argument('--weight_decay',type=float, default=0)
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
mkdir data
|
||||
cd data
|
||||
wget https://raw.githubusercontent.com/Impavidity/kim_cnn/master/data/stsa.fine.dev.tsv
|
||||
wget https://raw.githubusercontent.com/Impavidity/kim_cnn/master/data/stsa.fine.phrases.train.tsv
|
||||
wget https://raw.githubusercontent.com/Impavidity/kim_cnn/master/data/stsa.fine.test.tsv
|
||||
wget https://github.com/Impavidity/kim_cnn/raw/master/data/word2vec.sst-1.pt
|
||||
wget https://raw.githubusercontent.com/Impavidity/kim_cnn/master/vector_preprocess.py
|
||||
@@ -1,59 +0,0 @@
|
||||
import sys
|
||||
import random
|
||||
import numpy as np
|
||||
import torch
|
||||
from torchtext import data
|
||||
from args import get_args
|
||||
|
||||
from datasets.sst import SST1
|
||||
|
||||
args = get_args()
|
||||
torch.manual_seed(args.seed)
|
||||
if not args.cuda:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
print("Note: You are using GPU for training")
|
||||
torch.cuda.set_device(args.gpu)
|
||||
torch.cuda.manual_seed(args.seed)
|
||||
if torch.cuda.is_available() and not args.cuda:
|
||||
print("Warning: You have Cuda but do not use it. You are using CPU for training")
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
|
||||
if not args.trained_model:
|
||||
print("Error: You need to provide a option 'trained_model' to load the model")
|
||||
sys.exit(1)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu)
|
||||
|
||||
config = args
|
||||
config.target_class = train_iter.dataset.NUM_CLASSES
|
||||
config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab)
|
||||
config.embed_num = len(train_iter.dataset.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)
|
||||
|
||||
|
||||
def predict(dataset_iter, dataset_name):
|
||||
print("Dataset: {}".format(dataset_name))
|
||||
model.eval()
|
||||
dataset_iter.init_epoch()
|
||||
|
||||
n_correct = 0
|
||||
for data_batch_idx, data_batch in enumerate(dataset_iter):
|
||||
scores = model(data_batch)
|
||||
n_correct += (torch.max(scores, 1)[1].view(data_batch.label.size()).data == data_batch.label.data).sum()
|
||||
|
||||
print("no. correct {} out of {}".format(n_correct, len(dataset_iter.dataset.examples)))
|
||||
accuracy = 100. * n_correct / len(dataset_iter.dataset.examples)
|
||||
print("{} accuracy: {:8.6f}%".format(dataset_name, accuracy))
|
||||
|
||||
# Run the model on the dev set
|
||||
predict(dataset_iter=dev_iter, dataset_name="valid")
|
||||
|
||||
# Run the model on the test set
|
||||
predict(dataset_iter=test_iter, dataset_name="test")
|
||||
+4
-6
@@ -7,22 +7,20 @@ import torch.nn.functional as F
|
||||
class KimCNN(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(KimCNN, self).__init__()
|
||||
dataset = config.dataset
|
||||
output_channel = config.output_channel
|
||||
target_class = config.target_class
|
||||
words_num = config.words_num
|
||||
words_dim = config.words_dim
|
||||
embed_num = config.embed_num
|
||||
embed_dim = config.embed_dim
|
||||
self.mode = config.mode
|
||||
Ks = 3 # There are three conv net here
|
||||
Ks = 3 # There are three conv nets here
|
||||
if config.mode == 'multichannel':
|
||||
input_channel = 2
|
||||
else:
|
||||
input_channel = 1
|
||||
self.embed = nn.Embedding(words_num, words_dim)
|
||||
self.static_embed = nn.Embedding(embed_num, embed_dim)
|
||||
self.non_static_embed = nn.Embedding(embed_num, embed_dim)
|
||||
self.static_embed.weight.requires_grad = False
|
||||
self.static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=True)
|
||||
self.non_static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=False)
|
||||
|
||||
self.conv1 = nn.Conv2d(input_channel, output_channel, (3, words_dim), padding=(2,0))
|
||||
self.conv2 = nn.Conv2d(input_channel, output_channel, (4, words_dim), padding=(3,0))
|
||||
|
||||
@@ -1,135 +0,0 @@
|
||||
import time
|
||||
import os
|
||||
import random
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
|
||||
from datasets.sst import SST1
|
||||
from args import get_args
|
||||
from model import KimCNN
|
||||
|
||||
# Set default configuration in : args.py
|
||||
args = get_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(args.seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
if not args.cuda:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
print("Note: You are using GPU for training")
|
||||
torch.cuda.set_device(args.gpu)
|
||||
torch.cuda.manual_seed(args.seed)
|
||||
if torch.cuda.is_available() and not args.cuda:
|
||||
print("Warning: You have Cuda but not use it. You are using CPU for training.")
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
|
||||
# Set up the data for training SST-1
|
||||
if args.dataset == 'SST-1':
|
||||
train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu)
|
||||
|
||||
config = args
|
||||
config.target_class = train_iter.dataset.NUM_CLASSES
|
||||
config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab)
|
||||
config.embed_num = len(train_iter.dataset.TEXT_FIELD.vocab)
|
||||
|
||||
print("Dataset {} Mode {}".format(args.dataset, args.mode))
|
||||
print("VOCAB num",len(train_iter.dataset.TEXT_FIELD.vocab))
|
||||
print("LABEL.target_class:", train_iter.dataset.NUM_CLASSES)
|
||||
print("Train instance", len(train_iter.dataset))
|
||||
print("Dev instance", len(dev_iter.dataset))
|
||||
print("Test instance", len(test_iter.dataset))
|
||||
|
||||
if args.resume_snapshot:
|
||||
if args.cuda:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage)
|
||||
else:
|
||||
model = KimCNN(config)
|
||||
model.static_embed.weight.data.copy_(train_iter.dataset.TEXT_FIELD.vocab.vectors)
|
||||
model.non_static_embed.weight.data.copy_(train_iter.dataset.TEXT_FIELD.vocab.vectors)
|
||||
if args.cuda:
|
||||
model.cuda()
|
||||
print("Shift model to GPU")
|
||||
|
||||
|
||||
parameter = filter(lambda p: p.requires_grad, model.parameters())
|
||||
#for idx, p in enumerate(parameter):
|
||||
# print(idx, p)
|
||||
optimizer = torch.optim.Adadelta(parameter, lr=args.lr, weight_decay=args.weight_decay)
|
||||
criterion = nn.CrossEntropyLoss()
|
||||
early_stop = False
|
||||
best_dev_acc = 0
|
||||
iterations = 0
|
||||
iters_not_improved = 0
|
||||
epoch = 0
|
||||
start = time.time()
|
||||
header = ' Time Epoch Iteration Progress (%Epoch) Loss Dev/Loss Accuracy Dev/Accuracy'
|
||||
dev_log_template = ' '.join('{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{:8.6f},{:12.4f},{:12.4f}'.split(','))
|
||||
log_template = ' '.join('{:>6.0f},{:>5.0f},{:>9.0f},{:>5.0f}/{:<5.0f} {:>7.0f}%,{:>8.6f},{},{:12.4f},{}'.split(','))
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.save_path, args.dataset), exist_ok=True)
|
||||
print(header)
|
||||
|
||||
|
||||
while True:
|
||||
if early_stop:
|
||||
print("Early Stopping. Epoch: {}, Best Dev Acc: {}".format(epoch, best_dev_acc))
|
||||
break
|
||||
epoch += 1
|
||||
train_iter.init_epoch()
|
||||
n_correct, n_total = 0, 0
|
||||
|
||||
for batch_idx, batch in enumerate(train_iter):
|
||||
iterations += 1
|
||||
model.train()
|
||||
optimizer.zero_grad()
|
||||
scores = model(batch)
|
||||
n_correct += (torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum()
|
||||
n_total += batch.batch_size
|
||||
train_acc = 100. * n_correct / n_total
|
||||
|
||||
loss = criterion(scores, batch.label)
|
||||
loss.backward()
|
||||
|
||||
optimizer.step()
|
||||
|
||||
# Evaluate performance on validation set
|
||||
if iterations % args.dev_every == 1:
|
||||
# switch model into evalutaion mode
|
||||
model.eval()
|
||||
dev_iter.init_epoch()
|
||||
n_dev_correct = 0
|
||||
dev_losses = []
|
||||
for dev_batch_idx, dev_batch in enumerate(dev_iter):
|
||||
scores = model(dev_batch)
|
||||
n_dev_correct += (torch.max(scores, 1)[1].view(dev_batch.label.size()).data == dev_batch.label.data).sum()
|
||||
dev_loss = criterion(scores, dev_batch.label)
|
||||
dev_losses.append(dev_loss.item())
|
||||
dev_acc = 100. * n_dev_correct / len(dev_iter.dataset)
|
||||
print(dev_log_template.format(time.time() - start,
|
||||
epoch, iterations, 1 + batch_idx, len(train_iter),
|
||||
100. * (1 + batch_idx) / len(train_iter), loss.item(),
|
||||
sum(dev_losses) / len(dev_losses), train_acc, dev_acc))
|
||||
|
||||
# Update validation results
|
||||
if dev_acc > best_dev_acc:
|
||||
iters_not_improved = 0
|
||||
best_dev_acc = dev_acc
|
||||
snapshot_path = os.path.join(args.save_path, args.dataset, args.mode+'_best_model.pt')
|
||||
torch.save(model, snapshot_path)
|
||||
else:
|
||||
iters_not_improved += 1
|
||||
if iters_not_improved >= args.patience:
|
||||
early_stop = True
|
||||
break
|
||||
|
||||
if iterations % args.log_every == 1:
|
||||
# print progress message
|
||||
print(log_template.format(time.time() - start,
|
||||
epoch, iterations, 1 + batch_idx, len(train_iter),
|
||||
100. * (1 + batch_idx) / len(train_iter), loss.item(), ' ' * 8,
|
||||
n_correct / n_total * 100, ' ' * 12))
|
||||
@@ -1,21 +0,0 @@
|
||||
import re
|
||||
|
||||
|
||||
def clean_str(string):
|
||||
"""
|
||||
Tokenization/string cleaning for all datasets except for SST.
|
||||
"""
|
||||
string = re.sub(r"[^A-Za-z0-9(),!?\'\`]", " ", string)
|
||||
string = re.sub(r"\'s", " \'s", string)
|
||||
string = re.sub(r"\'ve", " \'ve", string)
|
||||
string = re.sub(r"n\'t", " n\'t", string)
|
||||
string = re.sub(r"\'re", " \'re", string)
|
||||
string = re.sub(r"\'d", " \'d", string)
|
||||
string = re.sub(r"\'ll", " \'ll", string)
|
||||
string = re.sub(r",", " , ", string)
|
||||
string = re.sub(r"!", " ! ", string)
|
||||
string = re.sub(r"\(", " ( ", string)
|
||||
string = re.sub(r"\)", " ) ", string)
|
||||
string = re.sub(r"\?", " ? ", string)
|
||||
string = re.sub(r"\s{2,}", " ", string)
|
||||
return string.lower().strip().split()
|
||||
Reference in New Issue
Block a user